use gam_linalg::faer_ndarray::fast_ata;
pub(crate) use super::tests_recovery_split_780::{
FdAnchorRegime, FdBranchRegime, FiniteDifferenceStratumCertificate,
certified_branch_stable_central_difference, certified_central_logdet_difference,
certified_fd_anchor, decisive_logit_homotopy, decisive_logit_pattern, diagonal_latent_cache,
fixed_state_logdet_sample, gamma_fd_tiny_fixture, rho_ladder_family,
rho_ladder_family_with_tolerance, smoothing_and_decisive_family, sparse_lift_ladder,
warmstart_test_objective, warmstart_test_objective_with_evaluator,
};
use super::*;
use approx::assert_abs_diff_eq;
use gam_terms::analytic_penalties::ARDPenalty;
use ndarray::{Array4, Array5, array};
pub(crate) fn assert_matrix_same_bits(left: &Array2<f64>, right: &Array2<f64>) {
assert_eq!(left.dim(), right.dim());
for ((row, col), &value) in left.indexed_iter() {
assert_eq!(
value.to_bits(),
right[[row, col]].to_bits(),
"matrix bits differ at ({row}, {col})"
);
}
}
pub(crate) fn assert_tensor3_same_bits(left: &Array3<f64>, right: &Array3<f64>) {
assert_eq!(left.dim(), right.dim());
for ((row, col, axis), &value) in left.indexed_iter() {
assert_eq!(
value.to_bits(),
right[[row, col, axis]].to_bits(),
"tensor bits differ at ({row}, {col}, {axis})"
);
}
}
pub(crate) fn assert_eta_one_parity(
evaluator: &dyn SaeBasisEvaluator,
coords: ArrayView2<'_, f64>,
expected_curved: usize,
) {
let (phi, jet) = evaluator.evaluate(coords).expect("base evaluate");
let eta = evaluator
.evaluate_phi_eta(coords, 1.0)
.expect("eta evaluate");
assert_matrix_same_bits(&eta.phi, &phi);
assert_tensor3_same_bits(&eta.jet, &jet);
assert_eq!(
eta.split.curved_cols.len(),
expected_curved,
"phi_eta split must report the expected curved-column count: got {} curved of {} total \
(base {}), expected {expected_curved}",
eta.split.curved_cols.len(),
eta.phi.ncols(),
eta.split.base_cols.len(),
);
for &col in &eta.split.base_cols {
for row in 0..phi.nrows() {
assert_eq!(eta.dphi_deta[[row, col]], 0.0);
for axis in 0..jet.shape()[2] {
assert_eq!(eta.djet_deta[[row, col, axis]], 0.0);
}
}
}
for &col in &eta.split.curved_cols {
for row in 0..phi.nrows() {
assert_eq!(
eta.dphi_deta[[row, col]].to_bits(),
phi[[row, col]].to_bits()
);
for axis in 0..jet.shape()[2] {
assert_eq!(
eta.djet_deta[[row, col, axis]].to_bits(),
jet[[row, col, axis]].to_bits()
);
}
}
}
}
#[test]
pub(crate) fn phi_eta_one_reproduces_current_atom_bases_bit_for_bit() {
let periodic_coords = array![[0.0_f64], [0.125], [0.4]];
let periodic = PeriodicHarmonicEvaluator::new(7).unwrap();
assert_eta_one_parity(&periodic, periodic_coords.view(), 4);
let raw_circle_coords = array![[0.0_f64], [0.3], [1.1]];
let raw_circle = RawPeriodicCircleEvaluator::new(1).unwrap();
assert_eta_one_parity(&raw_circle, raw_circle_coords.view(), 0);
let torus_coords = array![[0.0_f64, 0.2], [0.25, 0.5], [0.7, 0.9]];
let torus = TorusHarmonicEvaluator::new(2, 2).unwrap();
assert_eta_one_parity(&torus, torus_coords.view(), 20);
let sphere_coords = array![[0.0_f64, 0.0, 1.0], [0.6, -0.8, 0.0], [0.36, 0.48, 0.8]];
let sphere = AmbientSphereHarmonicEvaluator::new(2).unwrap();
assert_eta_one_parity(&sphere, sphere_coords.view(), 9);
let centers = array![
[-1.0_f64, -1.0],
[1.0, -1.0],
[-1.0, 1.0],
[1.0, 1.0],
[0.0, 0.0],
[0.5, -0.25]
];
let duchon_coords = array![[0.1_f64, 0.2], [0.4, -0.3], [-0.2, 0.7]];
let duchon = DuchonCoordinateEvaluator::new(centers, 2).unwrap();
let (duchon_phi, _) = duchon.evaluate(duchon_coords.view()).unwrap();
let duchon_poly = 3usize;
assert_eta_one_parity(
&duchon,
duchon_coords.view(),
duchon_phi.ncols() - duchon_poly,
);
let euclidean = EuclideanPatchEvaluator::new(2, 3).unwrap();
let total_cols = gam_terms::basis::monomial_exponents(2, 3).len();
let linear_cols = gam_terms::basis::monomial_exponents(2, 3)
.iter()
.filter(|alpha| alpha.iter().sum::<usize>() <= 1)
.count();
assert_eta_one_parity(&euclidean, duchon_coords.view(), total_cols - linear_cols);
}
pub(crate) fn trivial_k1_euclidean_term() -> SaeManifoldTerm {
let n = 4usize;
let p = 3usize;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"atom0",
SaeAtomBasisKind::EuclideanPatch,
1,
Array2::<f64>::ones((n, 2)),
Array3::<f64>::zeros((n, 2, 1)),
Array2::<f64>::zeros((2, p)),
Array2::<f64>::eye(2),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction");
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![Array2::<f64>::zeros((n, 1))],
vec![LatentManifold::Euclidean],
AssignmentMode::softmax(1.0),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
SaeManifoldTerm::new(vec![atom], assignment)
.expect("term fixture: every atom's row count matches the assignment's")
}
#[test]
pub(crate) fn criterion_gauge_deflation_count_bounded_flicker_reanchors_freely() {
let mut term = trivial_k1_euclidean_term();
term.record_criterion_gauge_deflation_count(150, true)
.unwrap();
let flicker = [
147usize, 150, 147, 150, 147, 150, 147, 150, 147, 150, 147, 150, 147, 150,
];
for &c in &flicker {
term.record_criterion_gauge_deflation_count(c, true)
.expect("a bounded low-amplitude flicker must re-anchor, never abort");
}
assert_eq!(
term.criterion_gauge_deflation_reanchors, 0,
"a flicker inside the relative jitter band charges no reversals"
);
assert_eq!(
term.expected_criterion_gauge_deflated_directions,
Some(150),
"the comparison re-anchors to the latest observed count"
);
let mut term2 = trivial_k1_euclidean_term();
term2
.record_criterion_gauge_deflation_count(150, true)
.unwrap();
let mut errored = false;
for &c in &[
40usize, 150, 40, 150, 40, 150, 40, 150, 40, 150, 40, 150, 40, 150,
] {
if term2
.record_criterion_gauge_deflation_count(c, true)
.is_err()
{
errored = true;
break;
}
}
assert!(
errored,
"a wide-amplitude oscillation must still exhaust the reversal budget"
);
}
#[test]
pub(crate) fn criterion_gauge_deflation_count_guard_reanchors_then_rejects_runaway() {
let mut term = trivial_k1_euclidean_term();
assert!(term.expected_criterion_gauge_deflated_directions.is_none());
term.record_criterion_gauge_deflation_count(60, true)
.unwrap();
assert_eq!(term.expected_criterion_gauge_deflated_directions, Some(60));
term.record_criterion_gauge_deflation_count(60, true)
.unwrap();
assert_eq!(term.expected_criterion_gauge_deflated_directions, Some(60));
for c in [50usize, 40, 33, 21, 12, 9, 6, 4, 3, 2] {
term.record_criterion_gauge_deflation_count(c, true)
.unwrap();
assert_eq!(term.expected_criterion_gauge_deflated_directions, Some(c));
}
assert_eq!(
term.criterion_gauge_deflation_reanchors, 0,
"monotone drift charges no reversals"
);
let mut last_ok = 2usize;
let oscillation = [9usize, 2, 9, 2, 9, 2, 9, 2, 9, 2, 9, 2, 9, 2];
let mut errored = false;
for &c in &oscillation {
match term.record_criterion_gauge_deflation_count(c, true) {
Ok(()) => {
last_ok = c;
}
Err(err) => {
assert!(
err.contains("not stabilizing") && err.contains("oscillated"),
"guard must report the oscillating quotient dimension explicitly; got: {err}"
);
assert_eq!(
term.expected_criterion_gauge_deflated_directions,
Some(last_ok)
);
errored = true;
break;
}
}
}
assert!(
errored,
"a sustained oscillation must exceed the reversal budget and error"
);
}
#[test]
pub(crate) fn curvature_homotopy_eta_inertness_probe_tracks_curved_columns() {
let term = trivial_k1_euclidean_term();
assert!(term.curvature_homotopy_eta_is_inert().unwrap());
let (term, _target, _rho) = small_two_atom_periodic_term();
assert!(term.curvature_homotopy_eta_is_inert().unwrap());
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(7).unwrap());
let coords = array![[0.05], [0.20], [0.55], [0.80], [0.35]];
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic7",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
Array2::<f64>::zeros((7, 1)),
Array2::<f64>::eye(7),
)
.unwrap()
.with_basis_evaluator(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((5, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
assert!(!term.curvature_homotopy_eta_is_inert().unwrap());
}
#[test]
pub(crate) fn linear_span_anchor_recovers_planted_two_plane_configuration() {
let n = 4usize;
let p = 4usize;
let phi = Array2::<f64>::ones((n, 2));
let jet = Array3::<f64>::zeros((n, 2, 1));
let decoder = Array2::<f64>::zeros((2, p));
let smooth = Array2::<f64>::eye(2);
let atoms = vec![
SaeManifoldAtom::new_with_provided_function_gram(
"plane0",
SaeAtomBasisKind::EuclideanPatch,
1,
phi.clone(),
jet.clone(),
decoder.clone(),
smooth.clone(),
)
.unwrap(),
SaeManifoldAtom::new_with_provided_function_gram(
"plane1",
SaeAtomBasisKind::EuclideanPatch,
1,
phi,
jet,
decoder,
smooth,
)
.unwrap(),
];
let coords = vec![Array2::<f64>::zeros((n, 1)), Array2::<f64>::zeros((n, 1))];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 2)),
coords,
vec![LatentManifold::Euclidean, LatentManifold::Euclidean],
AssignmentMode::softmax(1.0),
)
.unwrap();
let term = SaeManifoldTerm::new(atoms, assignment).unwrap();
let target = array![
[3.0_f64, 0.0, 0.0, 0.0],
[0.0, 2.0, 0.0, 0.0],
[0.0, 0.0, 1.5, 0.0],
[0.0, 0.0, 0.0, 1.0]
];
let anchor = linear_span_anchor(&term, target.view()).unwrap();
assert_eq!(anchor.atoms.len(), 2);
assert_abs_diff_eq!(anchor.residual_norm_sq, 0.0, epsilon = 1.0e-18);
let plane0 = array![[1.0_f64, 0.0], [0.0, 1.0], [0.0, 0.0], [0.0, 0.0]];
let plane1 = array![[0.0_f64, 0.0], [0.0, 0.0], [1.0, 0.0], [0.0, 1.0]];
let angle0 = anchor.atoms[0]
.frame
.max_principal_angle(plane0.view())
.unwrap();
let angle1 = anchor.atoms[1]
.frame
.max_principal_angle(plane1.view())
.unwrap();
assert_abs_diff_eq!(angle0, 0.0, epsilon = 1.0e-12);
assert_abs_diff_eq!(angle1, 0.0, epsilon = 1.0e-12);
}
pub(crate) fn circle_certificate_fixture(
radius: f64,
planes: &[(usize, usize)],
) -> SaeManifoldTerm {
let n = 16usize;
let p = 4usize;
let evaluator =
Arc::new(PeriodicHarmonicEvaluator::new(3).expect("3 is a positive harmonic order"));
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| row as f64 / n as f64);
let (phi, jet) = evaluator
.evaluate(coords.view())
.expect("the fixture coordinates lie in the evaluator's domain");
let mut atoms = Vec::with_capacity(planes.len());
let mut coord_blocks = Vec::with_capacity(planes.len());
for (atom_idx, &(axis_sin, axis_cos)) in planes.iter().enumerate() {
let mut decoder = Array2::<f64>::zeros((3, p));
decoder[[1, axis_sin]] = radius;
decoder[[2, axis_cos]] = radius;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
format!("circle_{atom_idx}"),
SaeAtomBasisKind::Periodic,
1,
phi.clone(),
jet.clone(),
decoder,
Array2::<f64>::eye(3),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_second_jet(evaluator.clone());
atoms.push(atom);
coord_blocks.push(coords.clone());
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, planes.len())),
coord_blocks,
vec![LatentManifold::Circle { period: 1.0 }; planes.len()],
AssignmentMode::softmax(1.0),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
let mut term = SaeManifoldTerm::new(atoms, assignment)
.expect("term fixture: every atom's row count matches the assignment's");
term.set_certificate_dispersion(1.0)
.expect("1.0 is a strictly positive certificate dispersion");
term
}
#[test]
pub(crate) fn dictionary_incoherence_report_orthogonal_frames_has_zero_mu_hat() {
let term = circle_certificate_fixture(4.0, &[(0, 1), (2, 3)]);
let report = dictionary_incoherence_report(&term).unwrap();
assert_abs_diff_eq!(report.mu_hat, 0.0, epsilon = 1.0e-12);
assert_eq!(report.per_atom_kappa_hat.len(), 2);
assert!(
report.snr_proxy > 1.0,
"fixture must cross the certificate SNR gate; got {}",
report.snr_proxy
);
match report.global_optimality {
GlobalOptimalityVerdict::CertifiedGlobal { margin } => assert!(margin > 0.0),
GlobalOptimalityVerdict::Uncertified { margin } => panic!(
"orthogonal frames with κ̂=0.25 and SNR>1 must certify; margin={margin:?}, note={}",
report.note
),
}
}
#[test]
pub(crate) fn dictionary_incoherence_report_coherent_frames_has_unit_mu_hat() {
let term = circle_certificate_fixture(2.0, &[(0, 1), (0, 1)]);
let report = dictionary_incoherence_report(&term).unwrap();
assert_abs_diff_eq!(report.mu_hat, 1.0, epsilon = 1.0e-12);
}
#[test]
pub(crate) fn dictionary_incoherence_report_circle_kappa_matches_inverse_radius() {
let radius = 2.5_f64;
let mut term = circle_certificate_fixture(radius, &[(0, 1)]);
term.set_certificate_dispersion(0.25).unwrap();
let report = dictionary_incoherence_report(&term).unwrap();
let kappa_hat = report.per_atom_kappa_hat[0]
.expect("a circle fixture of finite radius has a defined per-atom curvature");
assert_abs_diff_eq!(kappa_hat, 1.0 / radius, epsilon = 1.0e-10);
assert!(report.snr_proxy.is_finite() && report.snr_proxy > 0.0);
assert_abs_diff_eq!(report.mean_activity_floor, 1.0, epsilon = 1.0e-12);
assert_abs_diff_eq!(report.peak_activity_floor, 1.0, epsilon = 1.0e-12);
}
#[test]
pub(crate) fn k1_gate_modes_do_not_pin_assignment_to_one() {
let ordered_beta_bernoulli = SaeAssignment::from_blocks_with_mode(
array![[0.0]],
vec![array![[0.0]]],
AssignmentMode::ordered_beta_bernoulli(1.0, 1.0, false),
)
.unwrap();
let ordered_beta_bernoulli_gate = ordered_beta_bernoulli.try_assignments_row(0).unwrap()[0];
assert_abs_diff_eq!(ordered_beta_bernoulli_gate, 0.5, epsilon = 1e-9);
assert!(
(ordered_beta_bernoulli_gate - 1.0).abs() > 1e-6,
"K=1 ordered Beta--Bernoulli must not pin the gate to 1.0"
);
let jr = SaeAssignment::from_blocks_with_mode(
array![[-1.0]],
vec![array![[0.0]]],
AssignmentMode::threshold_gate(1.0, 0.0),
)
.unwrap();
assert_abs_diff_eq!(
jr.try_assignments_row(0).unwrap()[0],
gam_linalg::utils::stable_logistic(-1.0),
epsilon = 1e-12
);
let sm = SaeAssignment::from_blocks_with_mode(
Array2::<f64>::zeros((1, 1)),
vec![array![[0.0]]],
AssignmentMode::softmax(1.0),
)
.unwrap();
assert_abs_diff_eq!(sm.try_assignments_row(0).unwrap()[0], 1.0, epsilon = 1e-12);
}
#[test]
pub(crate) fn smooth_threshold_gate_is_centered_at_threshold() {
let threshold = 2.0;
let temperature = 1.0;
let logits = array![2.0 + 1e-6, 1.0];
let gates = threshold_gate_row(logits.view(), temperature, threshold);
assert_abs_diff_eq!(gates[0], 0.5, epsilon = 1e-3);
assert!(
gates[0] < 0.6,
"surrogate not centered at threshold: {}",
gates[0]
);
assert_abs_diff_eq!(
gates[1],
gam_linalg::utils::stable_logistic(-1.0),
epsilon = 1e-12
);
}
pub(crate) fn periodic_basis(coords: &Array2<f64>) -> (Array2<f64>, Array3<f64>) {
let n = coords.nrows();
let mut phi = Array2::<f64>::zeros((n, 3));
let mut jet = Array3::<f64>::zeros((n, 3, 1));
for row in 0..n {
let x = coords[[row, 0]].rem_euclid(1.0);
let angle = 2.0 * std::f64::consts::PI * x;
phi[[row, 0]] = 1.0;
phi[[row, 1]] = angle.sin();
phi[[row, 2]] = angle.cos();
jet[[row, 1, 0]] = 2.0 * std::f64::consts::PI * angle.cos();
jet[[row, 2, 0]] = -2.0 * std::f64::consts::PI * angle.sin();
}
(phi, jet)
}
#[test]
pub(crate) fn ard_axis_prior_periodic_is_continuous_across_cut() {
let alpha = 2.3_f64;
let period = 1.0_f64;
let eps = 1.0e-6;
let below = ArdAxisPrior::eval(alpha, period - eps, Some(period));
let above = ArdAxisPrior::eval(alpha, period + eps, Some(period));
let at_zero = ArdAxisPrior::eval(alpha, 0.0, Some(period));
let cont_tol = 10.0 * alpha * eps; assert!((below.value - above.value).abs() < cont_tol);
assert!((below.grad - above.grad).abs() < cont_tol);
assert!((below.hess - above.hess).abs() < cont_tol);
assert!(below.grad.abs() < cont_tol);
assert!(above.grad.abs() < cont_tol);
assert_abs_diff_eq!(below.value, at_zero.value, epsilon = 1.0e-9);
assert_abs_diff_eq!(at_zero.value, 0.0, epsilon = 1.0e-12);
assert_abs_diff_eq!(at_zero.grad, 0.0, epsilon = 1.0e-12);
assert_abs_diff_eq!(at_zero.hess, alpha, epsilon = 1.0e-12);
let sq_a = ArdAxisPrior::eval(1.0, 0.3, Some(period)).sq_equiv;
let sq_b = ArdAxisPrior::eval(5.0, 0.3, Some(period)).sq_equiv;
assert_abs_diff_eq!(sq_a, sq_b, epsilon = 1.0e-12);
let p = ArdAxisPrior::eval(alpha, 0.3, Some(period));
assert_abs_diff_eq!(0.5 * alpha * p.sq_equiv, p.value, epsilon = 1.0e-12);
}
#[test]
pub(crate) fn ard_axis_prior_value_grad_fd_consistent() {
let alpha = 1.7_f64;
let h = 1.0e-6;
for &period in &[None, Some(1.0_f64), Some(std::f64::consts::TAU)] {
for &t in &[-0.37_f64, 0.02, 0.49, 0.83, 0.999, 1.4] {
let p = ArdAxisPrior::eval(alpha, t, period);
let vp = ArdAxisPrior::eval(alpha, t + h, period).value;
let vm = ArdAxisPrior::eval(alpha, t - h, period).value;
let fd_grad = (vp - vm) / (2.0 * h);
assert_abs_diff_eq!(p.grad, fd_grad, epsilon = 1.0e-5);
let gp = ArdAxisPrior::eval(alpha, t + h, period).grad;
let gm = ArdAxisPrior::eval(alpha, t - h, period).grad;
let fd_hess = (gp - gm) / (2.0 * h);
assert_abs_diff_eq!(p.hess, fd_hess, epsilon = 1.0e-5);
}
}
}
#[test]
pub(crate) fn ard_axis_prior_tiny_energy_and_increment_are_resolved() {
let alpha = 1.7_f64;
let tiny = 1.0e-12;
let periodic = ArdAxisPrior::eval(alpha, tiny, Some(std::f64::consts::TAU));
let quadratic_limit = 0.5 * alpha * tiny * tiny;
assert!(
periodic.value > 0.0,
"nonzero periodic energy must not round to zero"
);
assert!(
(periodic.value - quadratic_limit).abs() <= 1.0e-15 * quadratic_limit,
"tiny periodic energy must retain its quadratic limit"
);
for &period in &[None, Some(std::f64::consts::TAU)] {
let from = 0.7_f64;
let to = from + tiny;
let delta = ArdAxisPrior::value_delta(alpha, from, to, period);
let first_order = ArdAxisPrior::eval(alpha, from, period).grad * (to - from);
assert!(delta.is_finite() && delta > 0.0);
assert!(
(delta - first_order).abs() <= 2.0e-12 * first_order.abs(),
"stable prior increment {delta:.17e} must agree with its local derivative {first_order:.17e}"
);
}
let period = 1.0_f64;
let across_cut = ArdAxisPrior::value_delta(alpha, period - tiny, tiny, Some(period));
assert!(
across_cut.abs() <= 1.0e-4 * alpha * tiny * tiny,
"equivalent points across the periodic cut must have the same energy"
);
}
#[test]
pub(crate) fn axis_periods_map_each_topology() {
assert_eq!(LatentManifold::Euclidean.axis_periods(), vec![None]);
assert_eq!(
LatentManifold::Circle { period: 1.0 }.axis_periods(),
vec![Some(1.0)]
);
let torus = LatentManifold::Product(vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
]);
assert_eq!(torus.axis_periods(), vec![Some(1.0), Some(1.0)]);
let sphere_chart = LatentManifold::Product(vec![
LatentManifold::Interval { lo: -1.0, hi: 1.0 },
LatentManifold::Circle {
period: std::f64::consts::TAU,
},
]);
assert_eq!(
sphere_chart.axis_periods(),
vec![None, Some(std::f64::consts::TAU)]
);
assert_eq!(
LatentManifold::Sphere { dim: 3 }.axis_periods(),
vec![None, None, None]
);
}
#[test]
pub(crate) fn ard_value_continuous_across_periodic_cut_d1() {
let coords0 = array![[0.999_f64]];
let (phi0, jet0) = periodic_basis(&coords0);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.2], [-0.3], [0.4]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((1, 1)),
vec![coords0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.1_f64]];
let alpha = 50.0_f64;
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![alpha.ln()]]);
let ard_before = term.loss(target.view(), &rho).unwrap().ard;
let q = term.assignment.row_block_dim();
let beta_dim = term.beta_dim();
let mut delta_ext = Array1::<f64>::zeros(q);
delta_ext[q - 1] = 0.002;
let delta_beta = Array1::<f64>::zeros(beta_dim);
term.apply_newton_step(delta_ext.view(), delta_beta.view(), 1.0)
.unwrap();
let wrapped = term.assignment.coords[0].row(0)[0];
assert!(
wrapped < 0.01,
"coordinate should have wrapped across the cut, got {wrapped}"
);
let ard_after = term.loss(target.view(), &rho).unwrap().ard;
assert!(
(ard_after - ard_before).abs() < 1.0e-2,
"periodic ARD jumped across the cut: before={ard_before}, after={ard_after}"
);
}
#[test]
pub(crate) fn penalized_objective_continuous_across_periodic_cut_with_registry_ard() {
let coords0 = array![[0.999_f64]];
let (phi0, jet0) = periodic_basis(&coords0);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.2], [-0.3], [0.4]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((1, 1)),
vec![coords0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.1_f64]];
let alpha = 50.0_f64;
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![alpha.ln()]]);
let coord = &term.assignment.coords[0];
let mut registry = AnalyticPenaltyRegistry::new();
let ard_pen = ARDPenalty::new(
PsiSlice::full(coord.len(), Some(coord.latent_dim())),
coord.latent_dim(),
);
registry.push(AnalyticPenaltyKind::Ard(Arc::new(ard_pen)));
let obj_before = term
.penalized_objective_total(target.view(), &rho, Some(®istry), 1.0)
.unwrap();
let q = term.assignment.row_block_dim();
let beta_dim = term.beta_dim();
let mut delta_ext = Array1::<f64>::zeros(q);
delta_ext[q - 1] = 0.002; let delta_beta = Array1::<f64>::zeros(beta_dim);
term.apply_newton_step(delta_ext.view(), delta_beta.view(), 1.0)
.unwrap();
let wrapped = term.assignment.coords[0].row(0)[0];
assert!(
wrapped < 0.01,
"coordinate should have wrapped across the cut, got {wrapped}"
);
let obj_after = term
.penalized_objective_total(target.view(), &rho, Some(®istry), 1.0)
.unwrap();
assert!(
(obj_after - obj_before).abs() < 1.0e-2,
"line-search objective jumped across the cut: before={obj_before}, after={obj_after}"
);
}
#[test]
pub(crate) fn scad_coord_penalty_inert_and_continuous_on_periodic_axis() {
use gam_terms::analytic_penalties::{PenaltyConcavity, ScadMcpPenalty};
let coords0 = array![[0.999_f64]];
let (phi0, jet0) = periodic_basis(&coords0);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.2], [-0.3], [0.4]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((1, 1)),
vec![coords0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.1_f64]];
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![0.0_f64]]);
let coord = &term.assignment.coords[0];
let mut registry = AnalyticPenaltyRegistry::new();
let scad = ScadMcpPenalty::new(
PsiSlice::full(coord.len(), Some(coord.latent_dim())),
5.0,
coord.n_obs(),
3.7,
1.0e-3,
PenaltyConcavity::Scad,
false,
)
.unwrap();
registry.push(AnalyticPenaltyKind::ScadMcp(Arc::new(scad)));
let with_scad = term
.penalized_objective_total(target.view(), &rho, Some(®istry), 1.0)
.unwrap();
let without = term
.penalized_objective_total(target.view(), &rho, None, 1.0)
.unwrap();
assert!(
(with_scad - without).abs() < 1.0e-12,
"SCAD coord penalty must be inert on a pure periodic axis: \
with={with_scad}, without={without}"
);
let obj_before = with_scad;
let q = term.assignment.row_block_dim();
let beta_dim = term.beta_dim();
let mut delta_ext = Array1::<f64>::zeros(q);
delta_ext[q - 1] = 0.002;
let delta_beta = Array1::<f64>::zeros(beta_dim);
term.apply_newton_step(delta_ext.view(), delta_beta.view(), 1.0)
.unwrap();
let wrapped = term.assignment.coords[0].row(0)[0];
assert!(
wrapped < 0.01,
"coordinate should have wrapped across the cut, got {wrapped}"
);
let obj_after = term
.penalized_objective_total(target.view(), &rho, Some(®istry), 1.0)
.unwrap();
assert!(
(obj_after - obj_before).abs() < 1.0e-2,
"SCAD line-search objective jumped across the periodic cut: \
before={obj_before}, after={obj_after}"
);
}
#[test]
pub(crate) fn scad_coord_penalty_active_on_euclidean_axis() {
let euclid = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((3, 1)),
vec![array![[0.5_f64], [-0.7], [1.3]]],
vec![LatentManifold::Euclidean],
AssignmentMode::softmax(0.7),
)
.unwrap();
assert!(
sae_coord_penalty_euclidean_restriction(&euclid.coords[0]).is_none(),
"Euclidean coord must not be restricted"
);
let circle = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((3, 1)),
vec![array![[0.1_f64], [0.4], [0.9]]],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let (axes, compacted) = sae_coord_penalty_euclidean_restriction(&circle.coords[0])
.expect("periodic coord must be restricted");
assert!(
axes.is_empty(),
"circle has no Euclidean axes, got {axes:?}"
);
assert_eq!(compacted.len(), 0, "compacted target must be empty");
}
#[test]
pub(crate) fn scad_no_origin_pinning_occupancy_on_circle() {
use gam_terms::analytic_penalties::{PenaltyConcavity, ScadMcpPenalty};
fn resultant_length(coords: &Array2<f64>) -> f64 {
let two_pi = 2.0 * std::f64::consts::PI;
let (mut cx, mut sy) = (0.0_f64, 0.0_f64);
for row in 0..coords.nrows() {
let a = two_pi * coords[[row, 0]];
cx += a.cos();
sy += a.sin();
}
let n = coords.nrows() as f64;
((cx / n).powi(2) + (sy / n).powi(2)).sqrt()
}
let scad_contribution = |coords: Array2<f64>| -> f64 {
let n = coords.nrows();
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.2_f64], [-0.3], [0.4]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let coord = &term.assignment.coords[0];
let mut registry = AnalyticPenaltyRegistry::new();
let scad = ScadMcpPenalty::new(
PsiSlice::full(coord.len(), Some(coord.latent_dim())),
5.0,
coord.n_obs(),
3.7,
1.0e-3,
PenaltyConcavity::Scad,
false,
)
.unwrap();
registry.push(AnalyticPenaltyKind::ScadMcp(Arc::new(scad)));
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![0.0_f64]]);
let target = Array2::<f64>::zeros((n, 1));
let with_scad = term
.penalized_objective_total(target.view(), &rho, Some(®istry), 1.0)
.unwrap();
let without = term
.penalized_objective_total(target.view(), &rho, None, 1.0)
.unwrap();
with_scad - without
};
let spread = array![[0.1_f64], [0.3], [0.5], [0.7], [0.9]];
let collapsed = array![[0.0_f64], [0.001], [0.0], [0.002], [0.001]];
assert!(
resultant_length(&spread) < 0.1,
"spread config should have near-uniform occupancy (R≈0), got {}",
resultant_length(&spread)
);
assert!(
resultant_length(&collapsed) > 0.99,
"collapsed config should have degenerate occupancy (R≈1), got {}",
resultant_length(&collapsed)
);
let c_spread = scad_contribution(spread);
let c_collapsed = scad_contribution(collapsed);
assert!(
c_spread.abs() < 1.0e-12,
"SCAD must add zero energy on a spread circle occupancy, got {c_spread}"
);
assert!(
c_collapsed.abs() < 1.0e-12,
"SCAD must add zero energy on a collapsed circle occupancy, got {c_collapsed}"
);
assert!(
(c_spread - c_collapsed).abs() < 1.0e-12,
"SCAD must not prefer origin-collapsed over spread occupancy \
(origin-pinning bias): spread={c_spread}, collapsed={c_collapsed}"
);
}
#[test]
pub(crate) fn shared_ard_flat_index_aliases_in_bounds_1026() {
let shared = SaeManifoldRho::new_shared_ard(0.0, 0.0, vec![array![0.1_f64], array![0.2_f64]]);
let shared_len = shared.to_flat().len();
assert_eq!(shared_len, 4, "shared flat len = 1+K+max_d");
assert_eq!(shared.ard_flat_index(0, 0), 3);
assert_eq!(
shared.ard_flat_index(1, 0),
3,
"both atoms' axis 0 alias the single shared coordinate"
);
assert!(
shared.ard_flat_index(1, 0) < shared_len,
"shared index must stay in bounds (the old per-atom walk went OOB)"
);
let per_atom = SaeManifoldRho::new(0.0, 0.0, vec![array![0.1_f64], array![0.2_f64]]);
assert_eq!(per_atom.to_flat().len(), 5, "per-atom flat len = 1+K+Σ d_k");
assert_eq!(per_atom.ard_flat_index(0, 0), 3);
assert_eq!(
per_atom.ard_flat_index(1, 0),
4,
"per-atom keeps unique coordinates (bit-for-bit the historical cursor)"
);
}
#[test]
pub(crate) fn periodic_ard_curvature_is_psd_in_assembled_htt() {
let coords0 = array![[0.40_f64], [0.60_f64]];
let (phi0, jet0) = periodic_basis(&coords0);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.2], [-0.3], [0.4]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((2, 1)),
vec![coords0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.1_f64], [0.2_f64]];
let alpha = 100.0_f64;
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![alpha.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
for (row_idx, row) in sys.rows.iter().enumerate() {
let d = row.htt.nrows();
for a in 0..d {
assert!(
row.htt[[a, a]] >= 0.0,
"row {row_idx} htt diagonal[{a}]={} must be PSD (von-Mises \
curvature clamped to its positive part)",
row.htt[[a, a]]
);
}
}
}
#[test]
pub(crate) fn compact_layout_riemannian_geometry_matches_dense_on_full_support() {
let coords_a = array![[0.12_f64], [0.37], [0.66], [0.91]];
let coords_b = array![[0.81_f64], [0.05], [0.48], [0.23]];
let (phi_a, jet_a) = periodic_basis(&coords_a);
let (phi_b, jet_b) = periodic_basis(&coords_b);
let atom_a = SaeManifoldAtom::new_with_provided_function_gram(
"circle_a",
SaeAtomBasisKind::Periodic,
1,
phi_a,
jet_a,
array![[0.20, -0.10], [-0.30, 0.25], [0.40, 0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let atom_b = SaeManifoldAtom::new_with_provided_function_gram(
"circle_b",
SaeAtomBasisKind::Periodic,
1,
phi_b,
jet_b,
array![[-0.15, 0.30], [0.22, -0.18], [0.33, 0.27]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let n = 4usize;
let logits = array![[0.4_f64, -0.2], [-0.1, 0.5], [0.3, 0.1], [-0.4, 0.2]];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![coords_a, coords_b],
vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
],
AssignmentMode::top_k_support(2),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom_a, atom_b], assignment).unwrap();
let target = array![
[0.10_f64, -0.05],
[0.20, 0.15],
[-0.12, 0.08],
[0.05, -0.20]
];
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![5.0_f64.ln()], array![5.0_f64.ln()]]);
let probe = SAE_DENSE_BETA_PENALTY_PROBE_MAX_DIM;
let dense = term
.assemble_arrow_schur_inner(target.view(), &rho, None, 1.0, probe, Some(None))
.unwrap();
let layout = SaeRowLayout::from_topk_gates(
&term.assignments_all_parallel(n).unwrap(),
2,
vec![1usize, 1usize],
term.assignment.coord_offsets(),
)
.unwrap();
let compact = term
.assemble_arrow_schur_inner(target.view(), &rho, None, 1.0, probe, Some(Some(layout)))
.unwrap();
assert_eq!(dense.rows.len(), compact.rows.len());
for (row_idx, (dr, cr)) in dense.rows.iter().zip(compact.rows.iter()).enumerate() {
assert_eq!(
dr.gt.len(),
cr.gt.len(),
"row {row_idx}: gt length mismatch (full-support compact must equal dense q)"
);
for a in 0..dr.gt.len() {
assert_abs_diff_eq!(dr.gt[a], cr.gt[a], epsilon = 1e-12);
}
assert_eq!(dr.htt.dim(), cr.htt.dim());
for a in 0..dr.htt.nrows() {
for b in 0..dr.htt.ncols() {
assert_abs_diff_eq!(dr.htt[[a, b]], cr.htt[[a, b]], epsilon = 1e-12);
}
}
assert_eq!(dr.htbeta.dim(), cr.htbeta.dim());
for a in 0..dr.htbeta.nrows() {
for b in 0..dr.htbeta.ncols() {
assert_abs_diff_eq!(dr.htbeta[[a, b]], cr.htbeta[[a, b]], epsilon = 1e-12);
}
}
}
let any_curvature = compact
.rows
.iter()
.any(|r| r.htt.iter().any(|&v| v.abs() > 1e-9));
assert!(
any_curvature,
"assembled compact htt is all-zero — the test data did not exercise curvature"
);
}
#[test]
pub(crate) fn topk_joint_chart_gauges_restrict_132_to_84_with_rank_two_2653() {
let n = 48usize;
let p = 6usize;
let coords0 = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| {
((row as f64 + 0.25) / n as f64).rem_euclid(1.0)
});
let coords1 = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| {
((row as f64 + 0.75) / n as f64).rem_euclid(1.0)
});
let (phi0, jet0) = periodic_basis(&coords0);
let (phi1, jet1) = periodic_basis(&coords1);
let mut decoder0 = Array2::<f64>::zeros((3, p));
decoder0[[1, 0]] = 1.0;
decoder0[[2, 1]] = 0.8;
let mut decoder1 = Array2::<f64>::zeros((3, p));
decoder1[[1, 3]] = 0.9;
decoder1[[2, 5]] = 1.05;
let atom0 = SaeManifoldAtom::new_with_provided_function_gram(
"circle0",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
decoder0,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let atom1 = SaeManifoldAtom::new_with_provided_function_gram(
"circle1",
SaeAtomBasisKind::Periodic,
1,
phi1,
jet1,
decoder1,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let logits = Array2::<f64>::from_shape_fn((n, 2), |(row, atom)| {
if row % 2 == atom { 1.0 } else { -1.0 }
});
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![coords0, coords1],
vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
],
AssignmentMode::top_k_support(1),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom0, atom1], assignment).unwrap();
let target = Array2::<f64>::zeros((n, p));
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::zeros(1), Array1::zeros(1)]);
let system = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("TopK circle system assembles");
let dense = term.dense_step_gauge_vectors().expect("dense phase gauges");
assert_eq!(dense.len(), 2, "one phase gauge per circle");
assert!(
dense.iter().all(|gauge| gauge.len() == 132),
"dense chart must be 48*2 coordinates plus 2*3*6 decoder variables"
);
assert_eq!(*system.row_offsets.last().unwrap(), 48);
assert_eq!(system.gb.len(), 36);
let compact = term
.joint_chart_gauge_basis_for_arrow_layout(
&system.row_offsets,
system.gb.len(),
"topk_joint_chart_gauges_restrict_132_to_84_with_rank_two_2653",
)
.expect("dense chart gauges map into the exact compact arrow chart");
assert_eq!(compact.len(), 2, "the mapped gauge span must retain rank two");
assert!(compact.iter().all(|gauge| gauge.len() == 84));
for i in 0..compact.len() {
for j in 0..compact.len() {
let expected = if i == j { 1.0 } else { 0.0 };
assert_abs_diff_eq!(compact[i].dot(&compact[j]), expected, epsilon = 1.0e-12);
}
}
}
#[test]
pub(crate) fn compact_mixed_dimensional_manifold_expands_euclidean_axes_2295() {
let circle_coords = array![[0.12_f64], [0.63]];
let (circle_phi, circle_jet) = periodic_basis(&circle_coords);
let circle_atom = SaeManifoldAtom::new_with_provided_function_gram(
"circle",
SaeAtomBasisKind::Periodic,
1,
circle_phi,
circle_jet,
Array2::<f64>::zeros((3, 2)),
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let plane_coords = array![[0.2_f64, -0.4], [0.7, 0.1]];
let plane_evaluator = Arc::new(EuclideanPatchEvaluator::new(2, 1).unwrap());
let (plane_phi, plane_jet) = plane_evaluator.evaluate(plane_coords.view()).unwrap();
let plane_atom = SaeManifoldAtom::new_with_provided_function_gram(
"plane",
SaeAtomBasisKind::EuclideanPatch,
2,
plane_phi,
plane_jet,
Array2::<f64>::zeros((3, 2)),
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(plane_evaluator);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((2, 2)),
vec![circle_coords, plane_coords],
vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Euclidean,
],
AssignmentMode::top_k_support(2),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![circle_atom, plane_atom], assignment).unwrap();
let assignments = vec![array![1.0_f64, 1.0], array![1.0, 1.0]];
let layout =
SaeRowLayout::from_topk_gates(&assignments, 2, vec![1, 2], term.assignment.coord_offsets())
.unwrap();
let (manifold, point) = term.compact_row_ext_manifold_and_point(0, &layout);
assert_eq!(point.len(), 3);
assert_eq!(manifold.ambient_dim(point.len()), point.len());
let gradient = array![0.3_f64, -0.2, 0.5];
let velocity = array![0.04_f64, -0.03, 0.02];
let euclidean_hessian = array![[2.0_f64, 0.1, 0.0], [0.1, 1.5, -0.2], [0.0, -0.2, 1.0]];
assert_eq!(manifold.project_point(point.view()).len(), 3);
assert_eq!(
manifold
.project_gradient_to_tangent(point.view(), gradient.view())
.len(),
3
);
assert_eq!(
manifold
.project_vector_to_gradient_tangent(point.view(), gradient.view(), velocity.view(),)
.len(),
3
);
assert_eq!(manifold.retract(point.view(), velocity.view()).len(), 3);
assert_eq!(
manifold
.riemannian_hessian_matrix(point.view(), gradient.view(), euclidean_hessian.view())
.dim(),
(3, 3)
);
}
#[test]
pub(crate) fn dense_assignment_budget_refuses_without_truncation() {
fn build_term(k: usize, curved: bool) -> SaeManifoldTerm {
let n = 4usize;
let coords: Array2<f64> = array![[0.12_f64], [0.37], [0.66], [0.91]];
let (phi, jet) = periodic_basis(&coords);
let atoms: Vec<SaeManifoldAtom> = (0..k)
.map(|j| {
SaeManifoldAtom::new_with_provided_function_gram(
format!("atom_{j}"),
SaeAtomBasisKind::Periodic,
1,
phi.clone(),
jet.clone(),
array![[0.20, -0.10], [-0.30, 0.25], [0.40, 0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator))
})
.collect();
let manifold = if curved {
LatentManifold::Circle { period: 1.0 }
} else {
LatentManifold::Euclidean
};
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, k)),
(0..k).map(|_| coords.clone()).collect(),
(0..k).map(|_| manifold.clone()).collect(),
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.unwrap();
SaeManifoldTerm::new(atoms, assignment).unwrap()
}
let k = 8usize;
let curved = build_term(k, true);
let euclidean = build_term(k, false);
let curved_required = curved.exact_dense_assignment_bytes();
let euclidean_required = euclidean.exact_dense_assignment_bytes();
assert_eq!(
curved_required, euclidean_required,
"exact dense memory accounting must not depend on coordinate geometry"
);
let too_small = curved_required.saturating_sub(1);
let error = curved
.require_exact_dense_assignment_budget(too_small)
.expect_err("an undersized budget must refuse the exact dense model");
assert!(error.contains("never silently truncated"));
curved
.require_exact_dense_assignment_budget(curved_required)
.expect("the exact required-byte boundary is admitted");
euclidean
.require_exact_dense_assignment_budget(euclidean_required)
.expect("the exact required-byte boundary is admitted");
}
#[derive(Debug)]
struct SnapshotLinearSecondJet2521;
impl SaeBasisEvaluator for SnapshotLinearSecondJet2521 {
fn second_jet_dyn(&self, coords: ArrayView2<'_, f64>) -> Option<Result<Array4<f64>, String>> {
Some(<Self as SaeBasisSecondJet>::second_jet(self, coords))
}
fn third_jet_dyn(&self, coords: ArrayView2<'_, f64>) -> Option<Result<Array5<f64>, String>> {
if coords.ncols() != 1 {
return Some(Err(format!(
"SnapshotLinearSecondJet2521: coordinate width {} != 1",
coords.ncols()
)));
}
None
}
fn evaluate(&self, coords: ArrayView2<'_, f64>) -> Result<(Array2<f64>, Array3<f64>), String> {
if coords.ncols() != 1 {
return Err(format!(
"SnapshotLinearSecondJet2521: coordinate width {} != 1",
coords.ncols()
));
}
let n = coords.nrows();
let mut basis = Array2::<f64>::zeros((n, 2));
let mut jet = Array3::<f64>::zeros((n, 2, 1));
for row in 0..n {
basis[[row, 0]] = 1.0;
basis[[row, 1]] = coords[[row, 0]];
jet[[row, 1, 0]] = 1.0;
}
Ok((basis, jet))
}
}
impl SaeBasisSecondJet for SnapshotLinearSecondJet2521 {
fn second_jet(&self, coords: ArrayView2<'_, f64>) -> Result<Array4<f64>, String> {
if coords.ncols() != 1 {
return Err(format!(
"SnapshotLinearSecondJet2521: coordinate width {} != 1",
coords.ncols()
));
}
Ok(Array4::<f64>::zeros((coords.nrows(), 2, 1, 1)))
}
}
fn structural_restore_fixture_2521() -> SaeManifoldTerm {
let coords = Array2::<f64>::zeros((4, 1));
let evaluator = Arc::new(SnapshotLinearSecondJet2521);
let (basis, jet) = evaluator
.evaluate(coords.view())
.expect("the fixture coordinates lie in the evaluator's domain");
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"rank-reduced-restore",
SaeAtomBasisKind::Linear,
1,
basis,
jet,
array![
[0.20, -0.10, 0.30, -0.40, 0.50, -0.60],
[0.70, 0.80, -0.90, 1.00, -1.10, 1.20]
],
Array2::<f64>::eye(2),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_second_jet(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((4, 1)),
vec![coords],
vec![LatentManifold::Euclidean],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
SaeManifoldTerm::new(vec![atom], assignment)
.expect("term fixture: every atom's row count matches the assignment's")
}
#[test]
pub(crate) fn snapshot_restore_round_trips_full_and_reduced_atom_topologies_2521() {
let mut term = structural_restore_fixture_2521();
term.atoms[0].decoder_frame = Some(
GrassmannFrame::from_orthonormal(
array![[1.0_f64], [0.0], [0.0], [0.0], [0.0], [0.0]],
array![1.0],
)
.unwrap(),
);
let target = Array2::<f64>::zeros((4, 6));
let rho = SaeManifoldRho::new(0.0, -6.0, vec![Array1::<f64>::zeros(1)]);
let full_snapshot = term.snapshot_mutable_state();
let full_basis = term.atoms[0].basis_values.clone();
let full_jet = term.atoms[0].basis_jacobian.clone();
let full_decoder = term.atoms[0].decoder_coefficients().clone();
let full_model = full_basis.dot(&full_decoder);
let full_loss_bits = term.loss(target.view(), &rho).unwrap().total().to_bits();
term.reduce_atoms_to_data_supported_rank().unwrap();
assert_eq!(term.atoms[0].decoder_coefficients().dim(), (1, 6));
assert_eq!(term.atoms[0].basis_values.dim(), (4, 1));
assert_eq!(term.atoms[0].basis_jacobian.dim(), (4, 1, 1));
assert_eq!(term.atoms[0].smooth_penalty().dim(), (1, 1));
assert!(term.atoms[0].decoder_frame.is_none());
assert_eq!(
term.atoms[0]
.reduced_column_map
.as_ref()
.map(|column_map| column_map.dim()),
Some((2, 1))
);
term.atoms[0].chart_canonicalized = true;
let reduced_snapshot = term.snapshot_mutable_state();
let reduced_basis = term.atoms[0].basis_values.clone();
let reduced_jet = term.atoms[0].basis_jacobian.clone();
let reduced_decoder = term.atoms[0].decoder_coefficients().clone();
let reduced_model = reduced_basis.dot(&reduced_decoder);
let reduced_loss_bits = term.loss(target.view(), &rho).unwrap().total().to_bits();
term.restore_mutable_state(&full_snapshot)
.expect("full topology restores over reduced topology");
assert!(term.matches_mutable_state(&full_snapshot));
assert_eq!(term.atoms[0].decoder_coefficients().dim(), (2, 6));
assert_eq!(term.atoms[0].basis_values.dim(), (4, 2));
assert_eq!(term.atoms[0].basis_jacobian.dim(), (4, 2, 1));
assert_eq!(term.atoms[0].smooth_penalty().dim(), (2, 2));
assert!(term.atoms[0].reduced_column_map.is_none());
let restored_frame = term.atoms[0]
.decoder_frame
.as_ref()
.expect("full snapshot restores its profiled decoder frame");
assert_eq!((restored_frame.output_dim(), restored_frame.rank()), (6, 1));
assert!(!term.atoms[0].chart_canonicalized);
assert_matrix_same_bits(&term.atoms[0].basis_values, &full_basis);
assert_tensor3_same_bits(&term.atoms[0].basis_jacobian, &full_jet);
assert_matrix_same_bits(term.atoms[0].decoder_coefficients(), &full_decoder);
assert_matrix_same_bits(
&term.atoms[0]
.basis_values
.dot(term.atoms[0].decoder_coefficients()),
&full_model,
);
assert_eq!(
term.loss(target.view(), &rho).unwrap().total().to_bits(),
full_loss_bits
);
term.restore_mutable_state(&reduced_snapshot)
.expect("reduced topology restores over full topology");
assert!(term.matches_mutable_state(&reduced_snapshot));
assert_eq!(term.atoms[0].decoder_coefficients().dim(), (1, 6));
assert_eq!(term.atoms[0].basis_values.dim(), (4, 1));
assert_eq!(term.atoms[0].basis_jacobian.dim(), (4, 1, 1));
assert_eq!(term.atoms[0].smooth_penalty().dim(), (1, 1));
assert_eq!(
term.atoms[0]
.reduced_column_map
.as_ref()
.map(|column_map| column_map.dim()),
Some((2, 1))
);
assert!(term.atoms[0].chart_canonicalized);
assert_matrix_same_bits(&term.atoms[0].basis_values, &reduced_basis);
assert_tensor3_same_bits(&term.atoms[0].basis_jacobian, &reduced_jet);
assert_matrix_same_bits(term.atoms[0].decoder_coefficients(), &reduced_decoder);
assert_matrix_same_bits(
&term.atoms[0]
.basis_values
.dot(term.atoms[0].decoder_coefficients()),
&reduced_model,
);
assert_eq!(
term.loss(target.view(), &rho).unwrap().total().to_bits(),
reduced_loss_bits
);
}
#[test]
pub(crate) fn snapshot_restore_refuses_incompatible_cardinality_atomically_2521() {
let mut term = structural_restore_fixture_2521();
let live = term.snapshot_mutable_state();
let mut incompatible = term.snapshot_mutable_state();
incompatible.atoms.clear();
let error = term
.restore_mutable_state(&incompatible)
.expect_err("missing atom must be a typed invariant refusal");
assert!(matches!(
error,
SaeMutableStateRestoreError::IncompatibleCardinality {
component: "atom",
expected,
observed,
} if expected == vec![1] && observed == vec![0]
));
assert!(
term.matches_mutable_state(&live),
"failed restore must not commit any prefix of the snapshot"
);
}
#[test]
pub(crate) fn snapshot_restore_round_trips_mutated_state() {
let coords0 = array![[0.05], [0.20], [0.55], [0.80]];
let (phi0, jet0) = periodic_basis(&coords0);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.2], [-0.3], [0.4]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((4, 1)),
vec![coords0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let snapshot = term.snapshot_mutable_state();
let pre_basis = term.atoms[0].basis_values.clone();
let pre_jet = term.atoms[0].basis_jacobian.clone();
let pre_decoder = term.atoms[0].decoder_coefficients().clone();
let pre_logits = term.assignment.logits.clone();
let pre_coords = term.assignment.coords[0].as_matrix();
let q = term.assignment.row_block_dim();
let beta_dim = term.beta_dim();
let delta_ext = Array1::<f64>::from_elem(4 * q, 0.3);
let delta_beta = Array1::<f64>::from_elem(beta_dim, -0.4);
term.apply_newton_step(delta_ext.view(), delta_beta.view(), 1.0)
.unwrap();
assert!(
(&term.atoms[0].basis_values - &pre_basis)
.mapv(f64::abs)
.sum()
> 1e-9
|| (term.atoms[0].decoder_coefficients() - &pre_decoder)
.mapv(f64::abs)
.sum()
> 1e-9,
"apply_newton_step did not perturb the snapshotted state"
);
term.restore_mutable_state(&snapshot)
.expect("differential restore rebuilds the basis");
assert_eq!(term.atoms[0].basis_values, pre_basis);
assert_eq!(term.atoms[0].basis_jacobian, pre_jet);
assert_eq!(term.atoms[0].decoder_coefficients(), pre_decoder);
assert_eq!(term.assignment.logits, pre_logits);
assert_eq!(term.assignment.coords[0].as_matrix(), pre_coords);
}
#[test]
pub(crate) fn ordered_beta_bernoulli_path_refreshes_periodic_basis_for_two_newton_iterations() {
let coords0 = array![[0.05], [0.20], [0.55], [0.80]];
let (phi0, jet0) = periodic_basis(&coords0);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.2], [-0.3], [0.4]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((4, 1)),
vec![coords0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, true),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.10], [0.05], [-0.15], [0.20]];
let mut rho = SaeManifoldRho::new(0.0, -6.0, vec![Array1::<f64>::zeros(1)]);
let loss0 = term.loss(target.view(), &rho).unwrap().total();
let basis0 = term.atoms[0].basis_values.clone();
let loss = term
.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 2, 0.05, 1.0e-3, 1.0e-3)
.unwrap();
assert!(loss.total().is_finite());
assert!(loss.total() <= loss0 + 1.0e-8);
assert!(
term.assignment.coords[0]
.as_flat()
.iter()
.all(|v| v.is_finite())
);
assert!(term.assignment.assignments().iter().all(|v| v.is_finite()));
let basis_delta = (&term.atoms[0].basis_values - &basis0).mapv(f64::abs).sum();
assert!(basis_delta > 1.0e-10);
}
#[test]
pub(crate) fn accepted_iterations_reuse_arrow_and_device_frame_allocations_with_fresh_content() {
let coords0 = array![[0.05], [0.20], [0.55], [0.80]];
let (phi0, jet0) = periodic_basis(&coords0);
let p = 16usize;
let mut decoder = Array2::<f64>::zeros((3, p));
decoder[[0, 0]] = 0.4;
decoder[[1, 0]] = -0.3;
decoder[[2, 0]] = 0.2;
decoder[[0, 1]] = -0.1;
decoder[[1, 1]] = 0.35;
decoder[[2, 1]] = 0.25;
let mut target = phi0.dot(&decoder);
for row in 0..target.nrows() {
target[[row, 0]] += 0.08 * (row as f64 + 0.5).sin();
target[[row, 1]] -= 0.06 * (row as f64 + 1.0).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"resident_periodic",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((4, 1)),
vec![coords0],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let mut rho = SaeManifoldRho::new(0.0, -4.0, vec![Array1::<f64>::zeros(1)]);
let decoder_before = term.atoms[0].decoder_coefficients().clone();
term.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 1, 0.05, 1.0e-3, 1.0e-3)
.expect("first accepted nonlinear iteration");
assert_ne!(
term.atoms[0].decoder_coefficients(),
decoder_before,
"first production call must accept a state-changing step"
);
let first_row = term
.arrow_assembly_workspace
.rows
.first()
.expect("driver returns row allocations to the workspace");
let first_row_htt_ptr = first_row.htt.as_ptr() as usize;
let first_row_htbeta_ptr = first_row.htbeta.as_ptr() as usize;
let first_gb_ptr = term.arrow_assembly_workspace.gb.as_ptr() as usize;
let first_device = term
.arrow_assembly_workspace
.device_sae_pcg
.as_ref()
.filter(|data| data.frame.is_some())
.expect("framed assembly returns its device descriptor");
let first_device_frame_ptr = Arc::as_ptr(first_device) as usize;
let first_device_frame_blocks_ptr = first_device
.frame
.as_ref()
.map_or(0, |frame| frame.frame_blocks.as_ptr() as usize);
let first_device_row_htbeta_ptr = first_device
.frame
.as_ref()
.and_then(|frame| frame.row_htbeta.first())
.map_or(0, |row| row.as_ptr() as usize);
let first_device_row_htbeta_bits: Vec<u64> = first_device
.frame
.as_ref()
.and_then(|frame| frame.row_htbeta.first())
.into_iter()
.flatten()
.map(|value| value.to_bits())
.collect();
let first_device_frame_block = first_device
.frame
.as_ref()
.and_then(|frame| frame.frame_blocks.first())
.expect("framed assembly retains a data-fit G tensor W block");
let first_device_frame_g_ptr = first_device_frame_block.g.as_ptr() as usize;
let first_device_frame_w_ptr = first_device_frame_block.w.as_ptr() as usize;
let first_device_frame_block_bits: Vec<u64> = first_device_frame_block
.g
.iter()
.chain(first_device_frame_block.w.iter())
.map(|value| value.to_bits())
.collect();
let first_numerical_bits: Vec<u64> = first_row
.htt
.iter()
.chain(first_row.htbeta.iter())
.chain(first_row.gt.iter())
.chain(term.arrow_assembly_workspace.gb.iter())
.map(|value| value.to_bits())
.collect();
term.atoms[0].decoder_coefficients_mut()[[0, 0]] += 0.02;
let decoder_before_second = term.atoms[0].decoder_coefficients().clone();
term.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 1, 0.05, 1.0e-3, 1.0e-3)
.expect("second accepted nonlinear iteration");
assert_ne!(
term.atoms[0].decoder_coefficients(),
decoder_before_second,
"second production call must accept a state-changing step"
);
let second_row = term
.arrow_assembly_workspace
.rows
.first()
.expect("driver returns reused row allocations to the workspace");
let second_device = term
.arrow_assembly_workspace
.device_sae_pcg
.as_ref()
.filter(|data| data.frame.is_some())
.expect("reused framed assembly returns its device descriptor");
let second_device_frame_ptr = Arc::as_ptr(second_device) as usize;
let second_device_frame_blocks_ptr = second_device
.frame
.as_ref()
.map_or(0, |frame| frame.frame_blocks.as_ptr() as usize);
let second_device_row_htbeta_ptr = second_device
.frame
.as_ref()
.and_then(|frame| frame.row_htbeta.first())
.map_or(0, |row| row.as_ptr() as usize);
let second_device_frame_block = second_device
.frame
.as_ref()
.and_then(|frame| frame.frame_blocks.first())
.expect("reused framed assembly retains its data-fit block");
let second_numerical_bits: Vec<u64> = second_row
.htt
.iter()
.chain(second_row.htbeta.iter())
.chain(second_row.gt.iter())
.chain(term.arrow_assembly_workspace.gb.iter())
.map(|value| value.to_bits())
.collect();
assert_ne!(first_row_htt_ptr, 0);
assert_eq!(first_row_htt_ptr, second_row.htt.as_ptr() as usize);
assert_ne!(first_row_htbeta_ptr, 0);
assert_eq!(first_row_htbeta_ptr, second_row.htbeta.as_ptr() as usize);
assert_ne!(first_gb_ptr, 0);
assert_eq!(
first_gb_ptr,
term.arrow_assembly_workspace.gb.as_ptr() as usize
);
assert_ne!(
first_device_frame_ptr, 0,
"framed assembly must install DeviceSaePcgData"
);
assert_eq!(first_device_frame_ptr, second_device_frame_ptr);
assert_ne!(first_device_frame_blocks_ptr, 0);
assert_eq!(
first_device_frame_blocks_ptr, second_device_frame_blocks_ptr,
"framed data-fit block vector must retain its allocation"
);
assert_ne!(first_device_row_htbeta_ptr, 0);
assert_eq!(
first_device_row_htbeta_ptr, second_device_row_htbeta_ptr,
"dominant framed row H_tbeta host slab must retain its allocation"
);
assert_ne!(first_device_frame_g_ptr, 0);
assert_eq!(
first_device_frame_g_ptr,
second_device_frame_block.g.as_ptr() as usize,
"framed data-fit G block must retain its allocation"
);
assert_ne!(first_device_frame_w_ptr, 0);
assert_eq!(
first_device_frame_w_ptr,
second_device_frame_block.w.as_ptr() as usize,
"framed output-factor W block must retain its allocation"
);
let second_device_row_htbeta_bits: Vec<u64> = second_device
.frame
.as_ref()
.and_then(|frame| frame.row_htbeta.first())
.into_iter()
.flatten()
.map(|value| value.to_bits())
.collect();
assert_ne!(
first_device_row_htbeta_bits, second_device_row_htbeta_bits,
"retained framed row H_tbeta slab must be numerically refreshed"
);
let second_device_frame_block_bits: Vec<u64> = second_device_frame_block
.g
.iter()
.chain(second_device_frame_block.w.iter())
.map(|value| value.to_bits())
.collect();
assert_ne!(
first_device_frame_block_bits, second_device_frame_block_bits,
"retained framed G tensor W block must be numerically refreshed"
);
assert_ne!(
first_numerical_bits, second_numerical_bits,
"accepted state change must refresh Hessian/gradient numerical content"
);
eprintln!(
"#1017 accepted-iteration residency telemetry: iterations=2 row_htt_ptr={} \
row_htbeta_ptr={} gb_ptr={} device_frame_ptr={} device_row_htbeta_ptr={} \
device_frame_blocks_ptr={} device_frame_g_ptr={} device_frame_w_ptr={} \
numerical_content_changed=true",
first_row_htt_ptr,
first_row_htbeta_ptr,
first_gb_ptr,
first_device_frame_ptr,
first_device_row_htbeta_ptr,
first_device_frame_blocks_ptr,
first_device_frame_g_ptr,
first_device_frame_w_ptr,
);
}
pub(crate) fn small_two_atom_periodic_term() -> (SaeManifoldTerm, Array2<f64>, SaeManifoldRho) {
let coords0 = array![[0.05], [0.20], [0.55], [0.80], [0.35]];
let coords1 = array![[0.15], [0.30], [0.65], [0.90], [0.45]];
let (phi0, jet0) = periodic_basis(&coords0);
let (phi1, jet1) = periodic_basis(&coords1);
let atom0 = SaeManifoldAtom::new_with_provided_function_gram(
"periodic0",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.25], [-0.35], [0.15]],
Array2::<f64>::eye(3),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let atom1 = SaeManifoldAtom::new_with_provided_function_gram(
"periodic1",
SaeAtomBasisKind::Periodic,
1,
phi1,
jet1,
array![[-0.10], [0.20], [0.30]],
Array2::<f64>::eye(3),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let logits = array![
[0.7, -0.2],
[0.1, 0.4],
[-0.3, 0.5],
[0.6, -0.1],
[0.2, 0.3]
];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![coords0, coords1],
vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
],
AssignmentMode::softmax(0.8),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
let term = SaeManifoldTerm::new(vec![atom0, atom1], assignment)
.expect("term fixture: every atom's row count matches the assignment's");
let target = array![[0.12], [-0.03], [0.08], [0.20], [-0.11]];
let rho = SaeManifoldRho::new(
(-0.3_f64).exp().ln(),
0.7_f64.ln(),
vec![array![0.9_f64.ln()], array![1.1_f64.ln()]],
);
(term, target, rho)
}
pub(crate) const FROZEN_INNER_STATE: usize = 0;
pub(crate) fn small_two_atom_periodic_term_at_shared_inner_state()
-> (SaeManifoldTerm, Array2<f64>, SaeManifoldRho) {
let (mut term, target, rho) = small_two_atom_periodic_term();
let drive_verdict = term.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
1,
0.25,
1.0e-4,
1.0e-4,
);
if let Err(err) = drive_verdict {
log::debug!(
"small_two_atom_periodic_term_at_shared_inner_state: the drive refused, as #2681 \
records it does on this fixture; the pinned state is whatever it left behind: {err:?}"
);
}
(term, target, rho)
}
#[test]
pub(crate) fn threshold_gate_fixed_logit_third_derivative_is_zero_bug4() {
use crate::manifold::arrow_solver::SaeLocalRowVar;
let (mut term, _target, rho) = small_two_atom_periodic_term();
term.assignment.mode = AssignmentMode::threshold_gate(1.0, 0.0);
term.assignment.ungated = vec![false, true];
for row in 0..term.n_obs() {
term.assignment.logits[[row, 0]] = 0.5;
term.assignment.logits[[row, 1]] = 0.5;
}
assert!(
term.assignment.logit_is_fixed(1) && !term.assignment.logit_is_fixed(0),
"atom 1 must be fixed (ungated), atom 0 free"
);
let threshold_strength = rho.lambda_sparse().unwrap();
let free = term.assignment_prior_hdiag_derivative_entry(
threshold_strength,
0,
0,
SaeLocalRowVar::Logit { atom: 0 },
None,
);
assert!(
free.abs() > 0.0,
"a FREE logit inside the band must carry a nonzero third derivative; got {free}"
);
let fixed = term.assignment_prior_hdiag_derivative_entry(
threshold_strength,
0,
1,
SaeLocalRowVar::Logit { atom: 1 },
None,
);
assert_eq!(
fixed, 0.0,
"a FIXED (ungated) logit third derivative must be zero; got {fixed}"
);
}
#[test]
pub(crate) fn per_atom_loao_ev_attributes_each_load_bearing_atom() {
let (term, _target, rho) = small_two_atom_periodic_term();
let target = term
.try_fitted_for_rho(&rho)
.expect("full reconstruction must assemble");
let ev_full = reconstruction_explained_variance(target.view(), target.view())
.expect("self-reconstruction EV defined");
assert!(
(ev_full - 1.0).abs() < 1e-12,
"target = full reconstruction ⇒ EV(full) = 1; got {ev_full}"
);
let dev = term
.per_atom_loao_explained_variance(target.view(), &rho)
.expect("LOAO EV must evaluate");
assert_eq!(dev.len(), term.k_atoms(), "one ΔEV per atom");
for (atom_idx, d) in dev.iter().enumerate() {
let d = d.unwrap_or_else(|| panic!("atom {atom_idx} ΔEV must be defined"));
assert!(
d > 1e-9,
"load-bearing atom {atom_idx} must earn positive training LOAO ΔEV; got {d:.3e}"
);
assert!(
d <= 1.0 + 1e-9,
"ΔEV for atom {atom_idx} cannot exceed EV(full)=1; got {d:.6e}"
);
}
let mut dead_term = term.clone();
dead_term.atoms[1].decoder_coefficients_mut().fill(0.0);
let dead_target = term
.try_fitted_for_rho(&rho)
.expect("reconstruction with the live atom-1 decoder");
let dead_dev = dead_term
.per_atom_loao_explained_variance(dead_target.view(), &rho)
.expect("LOAO EV must evaluate for the dead-atom term");
let d_dead = dead_dev[1].expect("dead atom ΔEV defined");
assert!(
d_dead.abs() < 1e-9,
"a zero-decoder atom carries no reconstruction ⇒ ΔEV ≈ 0; got {d_dead:.3e}"
);
}
fn collapse_rescue_term_and_target() -> (SaeManifoldTerm, Array2<f64>, SaeManifoldRho) {
let n = 6usize;
let coords = Array2::<f64>::from_elem((n, 1), 0.3);
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"collapsed_circle",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.1, -0.2], [0.05, 0.15], [-0.1, 0.08]],
Array2::<f64>::eye(3),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
let term = SaeManifoldTerm::new(vec![atom], assignment)
.expect("term fixture: every atom's row count matches the assignment's");
let s = [-0.5, -0.3, -0.1, 0.1, 0.3, 0.5];
let mu = [0.2, -0.1];
let d = [0.6, 0.8];
let mut target = Array2::<f64>::zeros((n, 2));
for row in 0..n {
target[[row, 0]] = mu[0] + s[row] * d[0];
target[[row, 1]] = mu[1] + s[row] * d[1];
}
let rho = SaeManifoldRho::new(0.02_f64.ln(), 1.0_f64.ln(), vec![array![0.0]]);
(term, target, rho)
}
#[test]
pub(crate) fn collapse_rescue_projection_matches_train_and_oos_and_refuses_targetless() {
let (mut term, target, rho) = collapse_rescue_term_and_target();
let report = term
.compute_hybrid_split_report(&rho, Some(target.view()))
.expect("hybrid split computes")
.expect("the collapsed d=1 atom presents a rescue verdict");
let rescue_image = report
.verdicts
.iter()
.find_map(|v| v.linear_image.clone())
.expect("the rescued slot carries a linear image");
assert!(
rescue_image.is_collapse_rescued() && rescue_image.v.is_some(),
"a collapse-rescued image must carry a projection direction v"
);
term.hybrid_split_report = Some(report);
let refusal = term
.try_fitted()
.expect_err("a rescued image cannot be reconstructed without its target");
assert!(
refusal.contains("requires try_fitted_target_aware"),
"unexpected target-less refusal: {refusal}"
);
let train_recon = term
.try_fitted_target_aware(target.view(), Some(&rho))
.expect("target-aware train reconstruction assembles");
let mut oos = term.clone();
oos.hybrid_split_report = None;
oos.set_hybrid_linear_images(vec![rescue_image.clone()])
.expect("trained rescue image attaches to the OOS term");
let oos_recon = oos
.try_fitted_target_aware(target.view(), Some(&rho))
.expect("target-aware OOS reconstruction assembles");
let max_gap = (&train_recon - &oos_recon)
.iter()
.fold(0.0_f64, |m, d| m.max(d.abs()));
assert!(
max_gap < 1e-10,
"train and OOS residual-projection reconstructions must be the SAME model \
within tol; max gap {max_gap:e}"
);
let ev = global_ev(target.view(), oos_recon.view());
assert!(
ev > 0.95,
"residual projection must recover the ramp; EV={ev:.4}"
);
}
#[test]
pub(crate) fn per_fit_config_isolates_barrier_and_ordered_beta_bernoulli_alpha() {
let (mut term_a, _t_a, rho_a) = small_two_atom_ordered_beta_bernoulli_term();
let (mut term_b, _t_b, rho_b) = small_two_atom_ordered_beta_bernoulli_term();
term_a.set_fit_config(SaeFitConfig {
separation_barrier_strength_override: Some(0.1),
ordered_beta_bernoulli_alpha_override: Some(0.2),
gpu_policy: gam_gpu::GpuPolicy::Off,
});
term_b.set_fit_config(SaeFitConfig {
separation_barrier_strength_override: Some(3.0),
ordered_beta_bernoulli_alpha_override: Some(5.0),
gpu_policy: gam_gpu::GpuPolicy::Required,
});
assert_eq!(
term_a.fit_config().ordered_beta_bernoulli_alpha_override,
Some(0.2)
);
assert_eq!(
term_b.fit_config().separation_barrier_strength_override,
Some(3.0)
);
assert_eq!(term_a.fit_config().gpu_policy, gam_gpu::GpuPolicy::Off);
assert_eq!(term_b.fit_config().gpu_policy, gam_gpu::GpuPolicy::Required);
assert_eq!(
term_a
.assignment
.resolved_ordered_beta_bernoulli_alpha(&rho_a),
Some(0.2)
);
assert_eq!(
term_b
.assignment
.resolved_ordered_beta_bernoulli_alpha(&rho_b),
Some(5.0)
);
assert_eq!(term_a.separation_barrier_strength(), 0.1);
assert_eq!(term_b.separation_barrier_strength(), 3.0);
term_a.set_fit_config(SaeFitConfig::default());
assert_eq!(
term_a
.assignment
.resolved_ordered_beta_bernoulli_alpha(&rho_a),
Some(1.0)
); assert_eq!(
term_b
.assignment
.resolved_ordered_beta_bernoulli_alpha(&rho_b),
Some(5.0)
);
}
#[test]
pub(crate) fn per_fit_barrier_isolated_under_concurrent_fits() {
let strengths = [0.125_f64, 7.5_f64];
let iters = 4000usize;
std::thread::scope(|scope| {
let handles: Vec<_> = strengths
.iter()
.map(|&mu| {
scope.spawn(move || {
let (mut term, _t, _rho) = small_two_atom_ordered_beta_bernoulli_term();
term.set_fit_config(SaeFitConfig {
separation_barrier_strength_override: Some(mu),
ordered_beta_bernoulli_alpha_override: None,
gpu_policy: gam_gpu::GpuPolicy::Off,
});
for _ in 0..iters {
assert_eq!(
term.separation_barrier_strength(),
mu,
"concurrent fit read a leaked barrier strength (expected {mu})"
);
}
mu
})
})
.collect();
for (handle, &mu) in handles.into_iter().zip(strengths.iter()) {
assert_eq!(handle.join().unwrap(), mu);
}
});
}
#[test]
pub(crate) fn assignment_logit_step_cap_bounds_single_iteration_gate_motion() {
let (mut term, _target, _rho) = small_two_atom_periodic_term();
let n = term.assignment.n_obs();
let q = term.assignment.row_block_dim();
let diff_before = term.assignment.logits[[0, 0]] - term.assignment.logits[[0, 1]];
let mut delta = Array1::<f64>::zeros(n * q);
delta[0] = 1.0e6;
let delta_beta = Array1::<f64>::zeros(term.beta_dim());
term.apply_newton_step(delta.view(), delta_beta.view(), 1.0)
.expect("step applies");
let cap = SAE_ASSIGNMENT_LOGIT_STEP_CAP_TAUS * term.assignment.mode.temperature();
let diff_after = term.assignment.logits[[0, 0]] - term.assignment.logits[[0, 1]];
assert!(
((diff_after - diff_before) - cap).abs() < 1.0e-9,
"a 1e6 raw logit delta must realise exactly the {cap}-cap, moved {}",
diff_after - diff_before
);
}
#[test]
pub(crate) fn active_mass_guard_reseeds_once_then_records_terminal_collapse() {
let (mut term, _target, _rho) = small_two_atom_periodic_term();
let n = term.assignment.n_obs();
let slam = |term: &mut SaeManifoldTerm| {
for row in 0..n {
term.assignment.logits[[row, 0]] = 0.0;
term.assignment.logits[[row, 1]] = -1.0e3;
}
};
slam(&mut term);
term.enforce_active_mass_guard(0, None).expect("guard runs");
assert_eq!(term.collapse_events().len(), 1);
let ev = term.collapse_events()[0];
assert_eq!(ev.atom, 1);
assert_eq!(ev.action, CollapseAction::Reseeded);
assert!(ev.max_active_mass < ev.floor);
let masses = term.assignment.assignments();
let max1 = (0..n).map(|r| masses[[r, 1]]).fold(0.0_f64, f64::max);
assert!(max1 > 1.0e-3_f64);
term.enforce_active_mass_guard(1, None).expect("guard runs");
assert_eq!(term.collapse_events().len(), 1);
slam(&mut term);
term.enforce_active_mass_guard(2, None).expect("guard runs");
term.enforce_active_mass_guard(3, None).expect("guard runs");
let terminals: Vec<_> = term
.collapse_events()
.iter()
.filter(|e| e.action == CollapseAction::Terminal)
.collect();
assert_eq!(terminals.len(), 1);
assert_eq!(terminals[0].atom, 1);
assert!(
term.collapse_events().iter().all(|e| e.atom == 1),
"the healthy atom must never be flagged"
);
}
#[test]
pub(crate) fn sae_rho_seed_dispersion_scaling_shifts_every_scale_coupled_axis() {
let rho = SaeManifoldRho::new(0.7_f64.ln(), 1.3_f64.ln(), vec![array![0.2, -0.4]]);
let dispersion = 0.05_f64 * 0.05;
let scaled = rho
.seed_scaled_by_dispersion_for_assignment(dispersion, AssignmentMode::softmax(1.0))
.unwrap();
let shift = dispersion.ln();
assert_abs_diff_eq!(
scaled.log_lambda_sparse,
rho.log_lambda_sparse + shift,
epsilon = 1.0e-14
);
assert_abs_diff_eq!(
scaled.log_lambda_smooth[0],
rho.log_lambda_smooth[0] + shift,
epsilon = 1.0e-14
);
assert_abs_diff_eq!(
scaled.log_ard[0][0],
rho.log_ard[0][0] + shift,
epsilon = 1.0e-14
);
assert_abs_diff_eq!(
scaled.log_ard[0][1],
rho.log_ard[0][1] + shift,
epsilon = 1.0e-14
);
for ordered_beta_bernoulli_mode in [
AssignmentMode::ordered_beta_bernoulli(1.0, 1.0, true),
AssignmentMode::ordered_beta_bernoulli(1.0, 1.0, false),
] {
let ordered_beta_bernoulli = rho
.seed_scaled_by_dispersion_for_assignment(dispersion, ordered_beta_bernoulli_mode)
.unwrap();
assert_abs_diff_eq!(
ordered_beta_bernoulli.log_lambda_sparse,
rho.log_lambda_sparse,
epsilon = 1.0e-14
);
assert_abs_diff_eq!(
ordered_beta_bernoulli.log_lambda_smooth[0],
rho.log_lambda_smooth[0],
epsilon = 1.0e-14
);
assert_abs_diff_eq!(
ordered_beta_bernoulli.log_ard[0][0],
rho.log_ard[0][0],
epsilon = 1.0e-14
);
assert_abs_diff_eq!(
ordered_beta_bernoulli.log_ard[0][1],
rho.log_ard[0][1],
epsilon = 1.0e-14
);
}
}
#[test]
pub(crate) fn fit_data_collapse_records_terminal_event_for_active_atom() {
let coords = array![[0.0], [0.25], [0.5], [0.75]];
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"circle",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
Array2::<f64>::zeros((3, 2)),
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((4, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[1.0, 0.0], [0.0, 1.0], [-1.0, 0.0], [0.0, -1.0]];
let rho = SaeManifoldRho::new(-0.3, 0.0, vec![array![0.0]]);
let recorded = term
.record_fit_data_collapse_if_needed(target.view(), &rho, 7)
.unwrap();
assert!(recorded);
let terminals: Vec<_> = term
.collapse_events()
.iter()
.filter(|event| event.action == CollapseAction::Terminal)
.collect();
assert_eq!(terminals.len(), 1);
assert_eq!(terminals[0].atom, 0);
assert_eq!(terminals[0].iteration, 7);
assert!(terminals[0].floor.is_finite() && terminals[0].floor >= 0.0);
}
pub(crate) fn deterministic_circle_noise(row: usize, col: usize) -> f64 {
let x = (row as f64 + 1.0) * 12.9898 + (col as f64 + 1.0) * 78.233;
(x.sin() * 43758.5453).sin()
}
pub(crate) fn planted_circle_data(n: usize, sigma: f64) -> Array2<f64> {
let mut z = Array2::<f64>::zeros((n, 2));
for row in 0..n {
let theta = std::f64::consts::TAU * row as f64 / n as f64;
z[[row, 0]] = theta.cos() + sigma * deterministic_circle_noise(row, 0);
z[[row, 1]] = theta.sin() + sigma * deterministic_circle_noise(row, 1);
}
z
}
pub(crate) fn planted_circle_embedded(n: usize, d_embed: usize, sigma: f64) -> Array2<f64> {
let mut frame = Array2::<f64>::zeros((2, d_embed));
for j in 0..d_embed {
frame[[0, j]] = deterministic_circle_noise(j, 0);
frame[[1, j]] = deterministic_circle_noise(j, 1);
}
for r in 0..2 {
let norm = (0..d_embed)
.map(|j| frame[[r, j]] * frame[[r, j]])
.sum::<f64>()
.sqrt();
for j in 0..d_embed {
frame[[r, j]] /= norm.max(1.0e-300);
}
}
let mut z = Array2::<f64>::zeros((n, d_embed));
for row in 0..n {
let theta = std::f64::consts::TAU * row as f64 / n as f64;
let (c, s) = (theta.cos(), theta.sin());
for j in 0..d_embed {
z[[row, j]] =
c * frame[[0, j]] + s * frame[[1, j]] + sigma * deterministic_circle_noise(row, j);
}
}
z
}
pub(crate) fn global_ev(target: ArrayView2<'_, f64>, fitted: ArrayView2<'_, f64>) -> f64 {
let (n, p) = target.dim();
let mut means = vec![0.0_f64; p];
for col in 0..p {
for row in 0..n {
means[col] += target[[row, col]];
}
means[col] /= n as f64;
}
let mut ssr = 0.0_f64;
let mut sst = 0.0_f64;
for row in 0..n {
for col in 0..p {
let r = target[[row, col]] - fitted[[row, col]];
ssr += r * r;
let centered = target[[row, col]] - means[col];
sst += centered * centered;
}
}
1.0 - ssr / sst.max(1.0e-300)
}
#[derive(Clone, Copy)]
pub(crate) enum PlantedCircleAssignmentMode {
Softmax,
OrderedBetaBernoulli,
}
impl PlantedCircleAssignmentMode {
pub(crate) fn label(self) -> &'static str {
match self {
Self::Softmax => "softmax",
Self::OrderedBetaBernoulli => "ordered_beta_bernoulli",
}
}
pub(crate) fn mode(self) -> AssignmentMode {
const TAU: f64 = 1.0;
const ALPHA: f64 = 1.0;
match self {
Self::Softmax => AssignmentMode::softmax(TAU),
Self::OrderedBetaBernoulli => AssignmentMode::ordered_beta_bernoulli(TAU, ALPHA, false),
}
}
pub(crate) fn seed_logit(self) -> f64 {
const TAU: f64 = 1.0;
match self {
Self::Softmax => 0.0,
Self::OrderedBetaBernoulli => 6.0 * TAU,
}
}
pub(crate) fn seed_gate(self) -> f64 {
match self {
Self::Softmax => 1.0,
Self::OrderedBetaBernoulli => 1.0 / (1.0 + (-6.0_f64).exp()),
}
}
}
pub(crate) fn planted_circle_seed_term(
z: ArrayView2<'_, f64>,
assignment_mode: PlantedCircleAssignmentMode,
) -> (SaeManifoldTerm, f64) {
let n = z.nrows();
let evaluator =
Arc::new(PeriodicHarmonicEvaluator::new(3).expect("3 is a positive harmonic order"));
let seed_coords = sae_pca_seed_initial_coords(z, &[SaeAtomBasisKind::Periodic], &[1])
.expect("one periodic atom of latent dimension 1 is a valid seed request");
let coords = seed_coords.slice(s![0, .., 0..1]).to_owned();
let (phi, jet) = evaluator
.evaluate(coords.view())
.expect("the fixture coordinates lie in the evaluator's domain");
let seed_gate = assignment_mode.seed_gate();
let gated_phi = &phi * seed_gate;
let mut xtx = fast_ata(&gated_phi);
for i in 0..xtx.nrows() {
xtx[[i, i]] += 1.0e-10;
}
let xtz = fast_atb(&gated_phi, &z.to_owned());
let decoder = xtx
.cholesky(Side::Lower)
.expect("the Gram is ridged by 1e-10 on its diagonal, so it is positive definite")
.solve_mat(&xtz);
let seed_fitted = gated_phi.dot(&decoder);
let mut rss = 0.0_f64;
for row in 0..n {
for col in 0..z.ncols() {
let r = z[[row, col]] - seed_fitted[[row, col]];
rss += r * r;
}
}
let seed_dispersion = (rss / (n * z.ncols()) as f64).max(1.0e-12);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"circle",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_evaluator(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::from_elem((n, 1), assignment_mode.seed_logit()),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
assignment_mode.mode(),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
(
SaeManifoldTerm::new(vec![atom], assignment)
.expect("term fixture: every atom's row count matches the assignment's"),
seed_dispersion,
)
}
#[test]
pub(crate) fn planted_circle_focus_1744() {
let n = 40usize;
let sigma = 0.05_f64;
let z = planted_circle_data(n, sigma);
let mut out = String::new();
for assignment_mode in [
PlantedCircleAssignmentMode::Softmax,
PlantedCircleAssignmentMode::OrderedBetaBernoulli,
] {
let label = assignment_mode.label();
let (term, seed_dispersion) = planted_circle_seed_term(z.view(), assignment_mode);
out.push_str(&format!(
"FOCUS1744 mode={label} seed_disp={seed_dispersion:.3e}\n"
));
for &sparse in &[-8.0_f64, 1.0] {
for &ard in &[-6.0_f64, -3.0, 0.0, 1.0] {
for &smooth in &[-8.0_f64, -5.0, -3.0, -1.0, 0.0, 1.0, 3.0] {
let mut t = term.clone();
let r = SaeManifoldRho::new(sparse, smooth, vec![array![ard]]);
match t.penalized_quasi_laplace_criterion_with_cache(
z.view(),
&r,
None,
60,
0.04,
1.0e-6,
1.0e-6,
) {
Ok(evaluated) => {
let ev = global_ev(z.view(), t.fitted().view());
out.push_str(&format!(
"FOCUS1744 mode={label} sparse={sparse} ard={ard} smooth={smooth} cost={:.4e} ev={ev:.4}\n",
evaluated.0
));
}
Err(err) => out.push_str(&format!(
"FOCUS1744 mode={label} sparse={sparse} ard={ard} smooth={smooth} ERR={err}\n"
)),
}
}
}
}
}
assert!(
out.contains("ev="),
"FOCUS1744: every (mode,sparse,ard,smooth) config errored — no fit produced a finite EV:\n{out}"
);
}
#[test]
pub(crate) fn planted_circle_ordered_beta_bernoulli_n40_sigma018_reaches_high_ev_1744() {
let assignment_mode = PlantedCircleAssignmentMode::OrderedBetaBernoulli;
let n = 40usize;
let sigma = 0.18_f64;
let z = planted_circle_data(n, sigma);
let (term, seed_dispersion) = planted_circle_seed_term(z.view(), assignment_mode);
let seed_ev = global_ev(z.view(), term.fitted().view());
let init_rho = SaeManifoldRho::new(0.02_f64.ln(), 1.0_f64.ln(), vec![array![0.0]])
.seed_scaled_by_dispersion_for_assignment(seed_dispersion, assignment_mode.mode())
.unwrap();
let init_rho_flat = init_rho.to_flat();
let n_params = init_rho_flat.len();
let mut objective =
SaeManifoldOuterObjective::new(term, z.clone(), None, init_rho, 50, 0.04, 1.0e-6, 1.0e-6);
let result = gam_solve::rho_optimizer::OuterProblem::new(n_params)
.with_initial_rho(init_rho_flat)
.run(&mut objective, "SAE planted circle #1744 focused")
.unwrap();
objective
.certify_outer_result(&result)
.expect("focused #1744 outer result must certify the installed state");
let fitted_result = objective.into_fitted().expect("outer fit was evaluated");
let rho = fitted_result.rho;
let ev = global_ev(z.view(), fitted_result.term.fitted().view());
assert!(
ev > 0.95,
"focused #1744 fixture (ordered_beta_bernoulli n={n} sigma={sigma}) seed_ev={seed_ev:.4} \
final_rho=({:.3},{:?},{:?}) EV={ev:.4} should exceed 0.95",
rho.log_lambda_sparse,
rho.log_lambda_smooth,
rho.log_ard
);
}
#[test]
pub(crate) fn planted_circle_noise_scale_sweep_reaches_high_ev_with_dimensionless_rho_seed() {
for assignment_mode in [
PlantedCircleAssignmentMode::Softmax,
PlantedCircleAssignmentMode::OrderedBetaBernoulli,
] {
let assignment_label = assignment_mode.label();
for &n in &[40usize, 250usize] {
for &sigma in &[0.02_f64, 0.05, 0.18] {
let z = planted_circle_data(n, sigma);
let (term, seed_dispersion) = planted_circle_seed_term(z.view(), assignment_mode);
let seed_ev = global_ev(z.view(), term.fitted().view());
let init_rho = SaeManifoldRho::new(0.02_f64.ln(), 1.0_f64.ln(), vec![array![0.0]])
.seed_scaled_by_dispersion_for_assignment(
seed_dispersion,
assignment_mode.mode(),
)
.unwrap();
let init_rho_flat = init_rho.to_flat();
let n_params = init_rho_flat.len();
let mut objective = SaeManifoldOuterObjective::new(
term,
z.clone(),
None,
init_rho,
50,
0.04,
1.0e-6,
1.0e-6,
);
let result = gam_solve::rho_optimizer::OuterProblem::new(n_params)
.with_initial_rho(init_rho_flat)
.run(&mut objective, "SAE planted circle dimensionless seed")
.unwrap();
objective
.certify_outer_result(&result)
.expect("dimensionless-seed outer result must certify the installed state");
let fitted_result = objective.into_fitted().expect("outer fit was evaluated");
let fitted_term = fitted_result.term;
let rho = fitted_result.rho;
let fitted = fitted_term.fitted();
let ev = global_ev(z.view(), fitted.view());
assert!(
ev > 0.95,
"planted circle assignment={assignment_label} n={n} sigma={sigma} seed_ev={seed_ev:.4} seed_phi={seed_dispersion:.3e} \
final_rho=({:.3}, {:?}, {:?}) EV={ev:.4} should exceed 0.95",
rho.log_lambda_sparse,
rho.log_lambda_smooth,
rho.log_ard
);
assert!(
fitted_term.collapse_events().is_empty(),
"healthy planted circle assignment={assignment_label} fit should not record collapse events: {:?}",
fitted_term.collapse_events()
);
}
}
}
}
#[test]
pub(crate) fn sae_value_probe_refusal_classification_is_inner_only() {
assert!(
SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(
"SaeManifoldTerm::penalized_quasi_laplace_criterion: inner solve did not converge at fixed ρ"
)
);
assert!(
SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(
"SaeManifoldTerm::penalized_quasi_laplace_criterion: undamped evidence factorization hit a non-PD per-row H_tt block before KKT stationarity and the refinement budget was exhausted"
)
);
assert!(
!SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(
"SaeManifoldTerm::penalized_quasi_laplace_criterion: arrow_log_det_from_cache returned None (undamped joint Hessian log-det unavailable for the Laplace normaliser)"
)
);
assert!(
!SaeManifoldOuterObjective::is_recoverable_value_probe_refusal(
"SaeManifoldTerm::penalized_quasi_laplace_criterion: row-gauge criterion deflation count re-anchored \
4 times within one optimization; the quotient dimension is not stabilizing"
)
);
}
#[test]
pub(crate) fn streaming_exact_reml_matches_full_batch_reml_small_sae() {
let (term0, target, rho) = small_two_atom_periodic_term_at_shared_inner_state();
let mut full = term0.clone();
let mut streaming = term0;
let (full_cost, full_loss, _cache) = full
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
FROZEN_INNER_STATE,
0.25,
1.0e-4,
1.0e-4,
)
.expect("dense REML must price the pinned inner state");
let (stream_cost, stream_loss) = streaming
.penalized_quasi_laplace_criterion_streaming_exact(
target.view(),
&rho,
None,
FROZEN_INNER_STATE,
0.25,
1.0e-4,
1.0e-4,
)
.expect("streaming REML must price the SAME pinned inner state");
let cost_gap = (stream_cost - full_cost).abs();
assert!(
cost_gap <= 1.0e-8,
"dense and streaming REML must price the SAME pinned inner state identically: \
stream_cost={stream_cost:?} full_cost={full_cost:?} |gap|={cost_gap:?} exceeds the \
1.0e-8 bound. This is the residual dense-vs-streaming EVIDENCE LOG-DETERMINANT desync \
left behind by #2509 Phase-2b, tracked as #2755 (measured 0.086061937491316, 3.86%). \
It is not a defect of this witness and it is NOT a bound to relax: the sibling gate \
`criterion_lane_gap_is_exactly_the_evidence_logdet_gap_2509` PASSES on this same \
fixture, which localises the WHOLE gap to the evidence log-determinant pair and proves \
the converged loss, the Occam term, the extra penalized energy and the MP rank charge \
all agree."
);
assert_abs_diff_eq!(stream_loss.total(), full_loss.total(), epsilon = 1.0e-8);
}
pub(crate) fn small_two_atom_ordered_beta_bernoulli_term()
-> (SaeManifoldTerm, Array2<f64>, SaeManifoldRho) {
let coords0 = array![[0.05], [0.20], [0.55], [0.80], [0.35]];
let coords1 = array![[0.15], [0.30], [0.65], [0.90], [0.45]];
let (phi0, jet0) = periodic_basis(&coords0);
let (phi1, jet1) = periodic_basis(&coords1);
let atom0 = SaeManifoldAtom::new_with_provided_function_gram(
"periodic0",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
array![[0.25], [-0.35], [0.15]],
Array2::<f64>::eye(3),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let atom1 = SaeManifoldAtom::new_with_provided_function_gram(
"periodic1",
SaeAtomBasisKind::Periodic,
1,
phi1,
jet1,
array![[-0.10], [0.20], [0.30]],
Array2::<f64>::eye(3),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let logits = array![
[0.7, -0.2],
[0.1, 0.4],
[-0.3, 0.5],
[0.6, -0.1],
[0.2, 0.3]
];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![coords0, coords1],
vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
],
AssignmentMode::ordered_beta_bernoulli(0.8, 1.0, false),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
let term = SaeManifoldTerm::new(vec![atom0, atom1], assignment)
.expect("term fixture: every atom's row count matches the assignment's");
let target = array![[0.12], [-0.03], [0.08], [0.20], [-0.11]];
let rho = SaeManifoldRho::new(
(-0.3_f64).exp().ln(),
0.7_f64.ln(),
vec![array![0.9_f64.ln()], array![1.1_f64.ln()]],
);
(term, target, rho)
}
#[test]
pub(crate) fn streaming_exact_laml_matches_full_batch_ordered_beta_bernoulli() {
let (term0, target, rho) = small_two_atom_ordered_beta_bernoulli_term();
let mut full = term0;
let (_full_cost, _full_loss, cache) = full
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
2,
0.25,
1.0e-4,
1.0e-4,
)
.expect("dense ordered Beta--Bernoulli criterion must evaluate");
let dense_logdet = arrow_log_det_from_cache(&cache).expect("dense log-det finite");
let stream_logdet = full
.streaming_exact_arrow_log_det(target.view(), &rho, None, None)
.expect("streaming ordered Beta--Bernoulli log-det must evaluate");
assert_abs_diff_eq!(stream_logdet, dense_logdet, epsilon = 1.0e-8);
}
#[test]
pub(crate) fn value_probe_refine_policy_ranks_same_criterion_as_full_policy() {
let (term0, target, rho) = small_two_atom_periodic_term();
let mut full = term0.clone();
let mut probe = term0;
let (full_cost, full_loss) = full
.penalized_quasi_laplace_criterion_with_refine_policy(
target.view(),
&rho,
None,
2,
0.25,
1.0e-4,
1.0e-4,
true,
)
.expect("full-budget criterion must converge on the small fixture");
let (probe_cost, probe_loss) = probe
.penalized_quasi_laplace_criterion_with_refine_policy(
target.view(),
&rho,
None,
2,
0.25,
1.0e-4,
1.0e-4,
false,
)
.expect("probe-budget criterion must converge on the small fixture");
assert_abs_diff_eq!(probe_cost, full_cost, epsilon = 1.0e-8);
assert_abs_diff_eq!(probe_loss.total(), full_loss.total(), epsilon = 1.0e-8);
}
#[test]
pub(crate) fn outer_value_and_ranking_lanes_share_pure_penalized_quasi_laplace_criterion() {
use gam_solve::rho_optimizer::{OuterEvalOrder, OuterObjective};
let rho_flat = warmstart_test_objective().baseline_rho.to_flat();
let mut grad_obj = warmstart_test_objective();
let grad_cost = grad_obj
.eval(&rho_flat)
.expect("gradient lane must converge on the warm-start fixture")
.cost;
assert!(
grad_obj.probe_telemetry.basin_envelope_evals > 0,
"a fresh dense gradient objective must seed and evaluate the authoritative envelope",
);
let mut ls_obj = warmstart_test_objective();
let ls_cost = ls_obj
.eval_with_order(&rho_flat, OuterEvalOrder::Value)
.expect("line-search probe must converge on the warm-start fixture")
.cost;
let mut adjacent_rho = rho_flat.clone();
adjacent_rho[0] = f64::from_bits(adjacent_rho[0].to_bits() ^ 1);
ls_obj
.eval_with_order(&adjacent_rho, OuterEvalOrder::Value)
.expect("adjacent line-search probe must seed a mismatched-rho handoff");
let preseeded_grad_cost = ls_obj
.eval(&rho_flat)
.expect("pre-seeded gradient lane must reconstruct the exact-rho handoff")
.cost;
let mut rank_obj = warmstart_test_objective();
let rank_cost = rank_obj
.eval_cost(&rho_flat)
.expect("ranking lane must converge on the warm-start fixture");
assert_abs_diff_eq!(ls_cost, grad_cost, epsilon = 1.0e-10);
assert_abs_diff_eq!(preseeded_grad_cost, grad_cost, epsilon = 1.0e-10);
assert_abs_diff_eq!(rank_cost, grad_cost, epsilon = 1.0e-10);
}
#[test]
pub(crate) fn refine_iteration_limit_probe_budget_never_extends() {
let probe_base = 16usize;
assert_eq!(
SaeManifoldTerm::refine_iteration_limit(
probe_base,
probe_base,
probe_base,
Some(1.0),
0.5,
true
),
probe_base
);
let accepted_base = 64usize;
let accepted_progress = 256usize;
assert_eq!(
SaeManifoldTerm::refine_iteration_limit(
accepted_base,
accepted_base,
accepted_progress,
Some(1.0),
0.5,
false
),
accepted_progress,
"accepted-point policy: a real residual drop (prev=Some(1.0), now=0.5) must extend the \
budget from base {accepted_base} to progress {accepted_progress}; got {}",
SaeManifoldTerm::refine_iteration_limit(
accepted_base,
accepted_base,
accepted_progress,
Some(1.0),
0.5,
false
),
);
assert_eq!(
SaeManifoldTerm::refine_iteration_limit(
accepted_base,
accepted_base,
accepted_progress,
Some(1.0),
1.0,
false
),
accepted_base
);
assert_eq!(
SaeManifoldTerm::refine_iteration_limit(
accepted_base - 1,
accepted_base,
accepted_progress,
None,
1.0e9,
false
),
accepted_base
);
}
#[test]
pub(crate) fn objective_stall_cannot_substitute_for_kkt_envelope_2253() {
let tolerance = 1.0e-4;
assert!(!SaeManifoldTerm::quasi_laplace_kkt_stationary(
2.0 * tolerance,
3.0 * tolerance,
tolerance,
));
assert!(SaeManifoldTerm::quasi_laplace_kkt_stationary(
tolerance,
3.0 * tolerance,
tolerance,
));
assert!(SaeManifoldTerm::quasi_laplace_kkt_stationary(
2.0 * tolerance,
tolerance,
tolerance,
));
assert!(!SaeManifoldTerm::quasi_laplace_kkt_stationary(
f64::NAN,
f64::INFINITY,
tolerance,
));
}
#[test]
pub(crate) fn reml_retries_refinement_after_non_pd_undamped_evidence_factor() {
let (mut term0, target, rho) = small_two_atom_periodic_term();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let cold_sys = term0
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let (.., cold_cache) = solve_arrow_newton_step_with_options(&cold_sys, 0.0, 0.0, &options)
.expect(
"cold undamped criterion factor must be spectrally conditioned (#1117), not refused",
);
let cold_deflated_rows = cold_cache
.deflation_row_spectra
.iter()
.filter(|spectrum| spectrum.is_some())
.count();
assert!(
cold_deflated_rows > 0 || cold_cache.gauge_deflated_directions > 0,
"fixture must start with a genuine non-PD evidence block that #1117 spectral \
unit-stiffness deflation had to condition; got no deflated row spectra and \
{} gauge directions",
cold_cache.gauge_deflated_directions,
);
let cold_grad_norm = SaeManifoldTerm::system_grad_norm_sq(&cold_sys).sqrt();
let (term0, target, rho) = small_two_atom_periodic_term_at_shared_inner_state();
let mut driven = term0.clone();
let driven_sys = driven
.assemble_arrow_schur(target.view(), &rho, None)
.expect("the driven state must assemble");
let driven_grad_norm = SaeManifoldTerm::system_grad_norm_sq(&driven_sys).sqrt();
assert!(
driven_grad_norm < cold_grad_norm,
"REML must refine PAST the cold non-PD undamped evidence factor: the driven state's KKT \
residual ‖g‖={driven_grad_norm:.6e} is not below the cold seed's {cold_grad_norm:.6e}, \
so no refinement survived the retry"
);
let mut full = term0.clone();
let mut streaming = term0;
let (full_cost, full_loss, cache) = full
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
FROZEN_INNER_STATE,
0.25,
1.0e-4,
1.0e-4,
)
.expect("dense REML must price the refined post-retry state");
let log_det = arrow_log_det_from_cache(&cache).expect("refined cache must carry log-det");
assert!(full_cost.is_finite());
assert!(full_loss.total().is_finite());
assert!(log_det.is_finite());
let (stream_cost, stream_loss) = streaming
.penalized_quasi_laplace_criterion_streaming_exact(
target.view(),
&rho,
None,
FROZEN_INNER_STATE,
0.25,
1.0e-4,
1.0e-4,
)
.expect("streaming REML must price the same refined post-retry state");
let cost_gap = (stream_cost - full_cost).abs();
assert!(
cost_gap <= 1.0e-8,
"dense and streaming REML must price the SAME pinned inner state identically: \
stream_cost={stream_cost:?} full_cost={full_cost:?} |gap|={cost_gap:?} exceeds the \
1.0e-8 bound. This is the residual dense-vs-streaming EVIDENCE LOG-DETERMINANT desync \
left behind by #2509 Phase-2b, tracked as #2755 (measured 0.086061937491316, 3.86%). \
It is not a defect of this witness and it is NOT a bound to relax: the sibling gate \
`criterion_lane_gap_is_exactly_the_evidence_logdet_gap_2509` PASSES on this same \
fixture, which localises the WHOLE gap to the evidence log-determinant pair and proves \
the converged loss, the Occam term, the extra penalized energy and the MP rank charge \
all agree."
);
assert_abs_diff_eq!(stream_loss.total(), full_loss.total(), epsilon = 1.0e-8);
}
#[test]
pub(crate) fn chunked_assembly_fold_is_bit_identical_1033() {
let (mut term_full, target, rho) = small_two_atom_periodic_term();
let mut term_chunked = term_full.clone();
term_full.assembly_chunk_override = None;
let sys_full = term_full
.assemble_arrow_schur(target.view(), &rho, None)
.expect("single-pass assembly must succeed");
term_chunked.assembly_chunk_override = Some(2);
let sys_chunked = term_chunked
.assemble_arrow_schur(target.view(), &rho, None)
.expect("chunked assembly must succeed");
assert_eq!(
sys_full.rows.len(),
sys_chunked.rows.len(),
"chunked and single-pass assemblies must yield the same row count"
);
assert_eq!(sys_full.gb.len(), sys_chunked.gb.len());
for (i, (gf, gc)) in sys_full.gb.iter().zip(sys_chunked.gb.iter()).enumerate() {
assert_eq!(
gf.to_bits(),
gc.to_bits(),
"sys.gb[{i}] must be bit-identical (single-pass {gf} vs chunked {gc})"
);
}
for (row, (rf, rc)) in sys_full
.rows
.iter()
.zip(sys_chunked.rows.iter())
.enumerate()
{
assert_eq!(rf.htt.dim(), rc.htt.dim(), "row {row} htt shape");
assert_eq!(rf.gt.len(), rc.gt.len(), "row {row} gt len");
for (a, (hf, hc)) in rf.htt.iter().zip(rc.htt.iter()).enumerate() {
assert_eq!(
hf.to_bits(),
hc.to_bits(),
"row {row} htt[{a}] must be bit-identical: {hf} vs {hc}"
);
}
for (a, (gf, gc)) in rf.gt.iter().zip(rc.gt.iter()).enumerate() {
assert_eq!(
gf.to_bits(),
gc.to_bits(),
"row {row} gt[{a}] must be bit-identical: {gf} vs {gc}"
);
}
}
}
#[test]
pub(crate) fn reconstruction_dispersion_uses_ard_shrunk_coordinate_edf() {
let n = 24usize;
let p = 2usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let alpha = 250.0_f64;
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![alpha.ln()]]);
let loss = term.loss(target.view(), &rho).unwrap();
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_delta_t, _delta_beta, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let dispersion = term
.reconstruction_dispersion(&loss, &cache, &rho, None)
.unwrap();
let smooth_edf: f64 = term
.decoder_smoothness_effective_dof_per_atom(&cache, &rho.lambda_smooth_vec().unwrap())
.unwrap()
.iter()
.sum();
let beta_edf = (term.beta_dim() as f64 - smooth_edf).max(0.0);
let traces = term.ard_shrinkage_traces(&cache).unwrap();
let coord_edf = super::construction_reconstruction::certified_ard_axis_edf(
n as f64,
alpha,
traces[0][0],
super::construction_reconstruction::undamped_row_curvature_scale(&cache),
0,
0,
)
.unwrap();
let rss = 2.0 * loss.data_fit;
let expected = rss / ((n * p) as f64 - beta_edf - coord_edf).max(1.0);
assert_abs_diff_eq!(dispersion, expected, epsilon = 1.0e-10);
let old_full_coordinate_edf = n as f64;
let old_full_coordinate_dispersion =
rss / ((n * p) as f64 - beta_edf - old_full_coordinate_edf).max(1.0);
let periods = term.assignment.coords[0].effective_axis_periods();
let coord_offsets = term.assignment.coord_offsets();
let inv_diag = cache.latent_block_inverse_diagonal().unwrap();
let mut unpenalized_rows = 0usize;
let mut assembled_shrinkage = 0.0_f64;
for row in 0..n {
let t = term.assignment.coords[0].row(row)[0];
let p_i = ArdAxisPrior::eval(alpha, t, periods[0]).psd_majorizer_hess();
if p_i == 0.0 {
unpenalized_rows += 1;
}
assembled_shrinkage += p_i * inv_diag[cache.row_offsets[row] + coord_offsets[0]];
}
assert_abs_diff_eq!(alpha * traces[0][0], assembled_shrinkage, epsilon = 1.0e-12);
assert_eq!(
unpenalized_rows,
n / 2,
"equispaced coordinates must split evenly across the von-Mises convex/concave arcs"
);
assert!(
coord_edf >= unpenalized_rows as f64,
"each slot the majorizer leaves unpenalized carries a full unit of \
freedom, so edf >= {unpenalized_rows}; got coord_edf={coord_edf}"
);
assert!(
coord_edf < old_full_coordinate_edf,
"the convex-arc slots ARE shrunk, so edf must sit strictly below the full \
coordinate count {old_full_coordinate_edf}; got coord_edf={coord_edf}"
);
assert!(
dispersion < old_full_coordinate_dispersion,
"φ̂ must use the ARD-shrunk coordinate edf, not the old full \
coordinate count: got {dispersion}, old formula {old_full_coordinate_dispersion}"
);
}
#[test]
fn matrix_free_smoothness_edf_from_probes_matches_dense_selected_inverse() {
let n = 24usize;
let p = 2usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.30, -0.10], [0.20, 0.40], [-0.35, 0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = Array2::from_shape_fn((n, p), |(row, col)| {
let x = (row as f64 + 0.5) / n as f64;
if col == 0 {
0.45 * (std::f64::consts::TAU * x).sin() + 0.07
} else {
-0.20 * (std::f64::consts::TAU * x).cos() + 0.03 * row as f64
}
});
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_delta_t, _delta_beta, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let lambda = rho.lambda_smooth_vec().unwrap();
let dense = term
.decoder_smoothness_effective_dof_per_atom(&cache, &lambda)
.unwrap();
let k = cache.k;
let sqrt_k = (k as f64).sqrt();
let probes: Vec<Array1<f64>> = (0..k)
.map(|j| {
let mut v = Array1::<f64>::zeros(k);
v[j] = sqrt_k;
v
})
.collect();
let sinv: Vec<Array1<f64>> = probes
.iter()
.map(|v| cache.schur_inverse_apply(v.view()).unwrap())
.collect();
let matrix_free = term
.decoder_smoothness_effective_dof_per_atom_from_probes(&probes, &sinv, &lambda)
.unwrap();
assert_eq!(dense.len(), matrix_free.len());
for (atom_idx, (d, mf)) in dense.iter().zip(&matrix_free).enumerate() {
assert_abs_diff_eq!(d, mf, epsilon = 1.0e-9);
assert!(
atom_idx < 1 || *mf >= 0.0,
"atom {atom_idx} edof must be a nonneg dof, got {mf}"
);
}
}
#[test]
fn matrix_free_ard_traces_from_probes_matches_dense_selected_inverse() {
let n = 24usize;
let p = 2usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.30, -0.10], [0.20, 0.40], [-0.35, 0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = Array2::from_shape_fn((n, p), |(row, col)| {
let x = (row as f64 + 0.5) / n as f64;
if col == 0 {
0.45 * (std::f64::consts::TAU * x).sin() + 0.07
} else {
-0.20 * (std::f64::consts::TAU * x).cos() + 0.03 * row as f64
}
});
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_delta_t, _delta_beta, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let dense = term.ard_inverse_traces(&cache).unwrap();
let k = cache.k;
let sqrt_k = (k as f64).sqrt();
let probes: Vec<Array1<f64>> = (0..k)
.map(|j| {
let mut v = Array1::<f64>::zeros(k);
v[j] = sqrt_k;
v
})
.collect();
let sinv: Vec<Array1<f64>> = probes
.iter()
.map(|v| cache.schur_inverse_apply(v.view()).unwrap())
.collect();
let matrix_free = term
.ard_inverse_traces_from_probes(&cache, &probes, &sinv)
.unwrap();
assert_eq!(dense.len(), matrix_free.len());
for (atom_idx, (d, mf)) in dense.iter().zip(&matrix_free).enumerate() {
assert_eq!(d.len(), mf.len());
for (axis, (dv, mv)) in d.iter().zip(mf.iter()).enumerate() {
assert_abs_diff_eq!(dv, mv, epsilon = 1.0e-9);
assert!(
*mv >= -1.0e-12,
"atom {atom_idx} axis {axis} posterior-variance trace must be nonneg, got {mv}"
);
}
}
}
#[test]
fn matrix_free_ard_logdet_hessian_trace_from_probes_matches_dense() {
let n = 24usize;
let p = 2usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.30, -0.10], [0.20, 0.40], [-0.35, 0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = Array2::from_shape_fn((n, p), |(row, col)| {
let x = (row as f64 + 0.5) / n as f64;
if col == 0 {
0.45 * (std::f64::consts::TAU * x).sin() + 0.07
} else {
-0.20 * (std::f64::consts::TAU * x).cos() + 0.03 * row as f64
}
});
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_delta_t, _delta_beta, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let solver = DeflatedArrowSolver::plain(&cache);
let dense = term
.ard_log_precision_hessian_trace(&rho, &cache, &solver, EvidenceOperator::Majorizer)
.unwrap();
let k = cache.k;
let sqrt_k = (k as f64).sqrt();
let probes: Vec<Array1<f64>> = (0..k)
.map(|j| {
let mut v = Array1::<f64>::zeros(k);
v[j] = sqrt_k;
v
})
.collect();
let sinv: Vec<Array1<f64>> = probes
.iter()
.map(|v| cache.schur_inverse_apply(v.view()).unwrap())
.collect();
let matrix_free = term
.ard_log_precision_hessian_trace_from_probes(
&rho,
&cache,
&probes,
&sinv,
EvidenceOperator::Majorizer,
)
.unwrap();
assert_eq!(dense.len(), matrix_free.len());
for (d, mf) in dense.iter().zip(&matrix_free) {
assert_eq!(d.len(), mf.len());
for (dv, mv) in d.iter().zip(mf.iter()) {
assert_abs_diff_eq!(dv, mv, epsilon = 1.0e-9);
}
}
}
#[test]
fn exact_a_route_gap_is_two_coordinates_with_two_causes_2515() {
let n = 24usize;
let p = 2usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
assert_eq!(decoder.ncols(), p);
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_delta_t, _delta_beta, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let loss = term.loss(target.view(), &rho).unwrap();
let delta_by_flat = term
.exact_stationarity_penalty_derivative_delta_by_flat(&rho, &cache)
.expect("the exact-A penalty derivative delta must be assemblable at this state");
for atom_idx in 0..rho.k_atoms() {
let smooth = rho.smooth_flat_index(atom_idx);
assert!(
!delta_by_flat.contains_key(&smooth),
"#2515: ΔC must not depend on λ_smooth (atom {atom_idx}, flat coordinate \
{smooth}), so ∂A/∂λ_smooth == ∂B/∂λ_smooth and the smooth-coordinate \
desync is a wrong-inverse gap, not a wrong-operator one. Keys present: {:?}",
delta_by_flat.keys().collect::<Vec<_>>()
);
}
let ard_flat = rho.ard_flat_index(0, 0);
assert!(
delta_by_flat.contains_key(&ard_flat),
"#2515: this fixture must carry a live ∂ΔC/∂log α on the ARD coordinate \
{ard_flat} (α=250 on a unit-period Circle puts cos κt below zero over a third \
of the rows), or it cannot distinguish the wrong-operator cause from the \
wrong-inverse one. Keys present: {:?}",
delta_by_flat.keys().collect::<Vec<_>>()
);
let solver = DeflatedArrowSolver::plain(&cache);
let dense = term
.analytic_outer_rho_gradient_components(target.view(), &rho, &loss, &cache, &solver)
.unwrap();
let k = cache.k;
let sqrt_k = (k as f64).sqrt();
let probes: Vec<Array1<f64>> = (0..k)
.map(|j| {
let mut v = Array1::<f64>::zeros(k);
v[j] = sqrt_k;
v
})
.collect();
let sinv: Vec<Array1<f64>> = probes
.iter()
.map(|v| cache.schur_inverse_apply(v.view()).unwrap())
.collect();
let bundled = term
.analytic_outer_rho_gradient_components_with_bundle(
target.view(),
&rho,
&loss,
&cache,
&solver,
Some(BundleEvidenceGeometry {
operator: EvidenceOperator::Majorizer,
cache: &cache,
probes: &probes,
sinv: &sinv,
}),
None,
)
.unwrap();
for i in 0..dense.logdet_trace.len() {
let role = if Some(i) == rho.sparse_flat_index() {
"assignment-log-strength".to_string()
} else if i >= rho.smooth_flat_start() && i < rho.smooth_flat_start() + rho.k_atoms() {
format!("smooth-atom-{} (wrong-inverse only)", i - rho.smooth_flat_start())
} else if delta_by_flat.contains_key(&i) {
format!("ard-flat-{i} (wrong-inverse AND wrong-operator)")
} else {
format!("ard-flat-{i}")
};
println!(
"[#2515 ROUTE-GAP] coord {i} ({role}) dense_A={:.17e} bundle_on_B={:.17e} \
|Δ|={:.6e}",
dense.logdet_trace[i],
bundled.logdet_trace[i],
(dense.logdet_trace[i] - bundled.logdet_trace[i]).abs(),
);
}
term.exact_a_evidence_system(target.view(), &rho, &sys)
.expect("the assembled exact-A evidence system must be constructible");
let dense_grad = dense.gradient();
let bundled_grad = bundled.gradient();
assert_eq!(dense_grad.len(), bundled_grad.len());
assert!(
dense_grad
.iter()
.chain(bundled_grad.iter())
.all(|v| v.is_finite()),
"#2515: both routes must produce a finite gradient: dense={dense_grad:?} \
bundled={bundled_grad:?}"
);
}
pub(crate) const RETIRED_ABSOLUTE_PD_FLOOR_REL_2673: f64 = 1.0e-9;
pub(crate) struct TwoFloorOverlap2673 {
pub value_gauge_gradient_resolved: usize,
pub value_priced_gradient_projected: usize,
pub gauge_band_population: usize,
pub priced_population: usize,
pub min_abs_lambda: f64,
pub max_abs_lambda: f64,
pub min_abs_lambda_priced: f64,
pub min_vbv: f64,
pub max_vbv: f64,
pub rows: Vec<String>,
}
pub(crate) fn classify_two_floor_overlap_2673(
a_eigs: ndarray::ArrayView1<'_, f64>,
a_vecs: ndarray::ArrayView2<'_, f64>,
b: &Array2<f64>,
floor: f64,
mu_floor: f64,
) -> TwoFloorOverlap2673 {
let mut out = TwoFloorOverlap2673 {
value_gauge_gradient_resolved: 0,
value_priced_gradient_projected: 0,
gauge_band_population: 0,
priced_population: 0,
min_abs_lambda: f64::INFINITY,
max_abs_lambda: 0.0,
min_abs_lambda_priced: f64::INFINITY,
min_vbv: f64::INFINITY,
max_vbv: f64::NEG_INFINITY,
rows: Vec::new(),
};
for idx in 0..a_eigs.len() {
let v = a_vecs.column(idx);
let vav = a_eigs[idx];
let vbv = v.dot(&b.dot(&v));
let mu = if vbv.abs() > 0.0 { vav / vbv } else { f64::NAN };
let value_calls_gauge = vav.abs() <= floor;
let gradient_calls_resolved = mu.is_finite() && mu.abs() >= mu_floor;
out.min_abs_lambda = out.min_abs_lambda.min(vav.abs());
out.max_abs_lambda = out.max_abs_lambda.max(vav.abs());
out.min_vbv = out.min_vbv.min(vbv);
out.max_vbv = out.max_vbv.max(vbv);
if value_calls_gauge {
out.gauge_band_population += 1;
} else {
out.priced_population += 1;
out.min_abs_lambda_priced = out.min_abs_lambda_priced.min(vav.abs());
}
if value_calls_gauge && gradient_calls_resolved {
out.value_gauge_gradient_resolved += 1;
out.rows.push(format!(
" direction {idx}: VALUE=gauge (|λ|={:.6e} <= floor {floor:.6e}) but \
GRADIENT=resolved (|μ|={:.6e} >= {mu_floor:.6e})",
vav.abs(),
mu.abs()
));
}
if !value_calls_gauge && mu.is_finite() && mu.abs() < mu_floor {
out.value_priced_gradient_projected += 1;
out.rows.push(format!(
" direction {idx}: VALUE=priced (|λ|={:.6e} > floor {floor:.6e}) but \
GRADIENT=projected out (|μ|={:.6e} < {mu_floor:.6e})",
vav.abs(),
mu.abs()
));
}
}
out
}
#[test]
fn two_floor_overlap_predicate_detects_both_crossings_2673() {
let mu_floor = f64::EPSILON.sqrt();
let lambdas = [5.0e-10_f64, 1.0, 0.5];
let max_eig = lambdas.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let floor = RETIRED_ABSOLUTE_PD_FLOOR_REL_2673 * max_eig.max(1.0);
assert!(
(floor - 1.0e-9).abs() <= 1.0e-24,
"#2673 control: floor must be 1e-9 here, got {floor:.6e}"
);
let vbvs = [1.0e-2_f64, 1.0e10, 1.0];
let a_eigs = Array1::from_vec(lambdas.to_vec());
let a_vecs = Array2::<f64>::eye(3);
let b = Array2::from_diag(&Array1::from_vec(vbvs.to_vec()));
let report = classify_two_floor_overlap_2673(a_eigs.view(), a_vecs.view(), &b, floor, mu_floor);
println!(
"[#2673 CONTROL] gauge&resolved={} priced&projected={} band_pop={} priced_pop={}",
report.value_gauge_gradient_resolved,
report.value_priced_gradient_projected,
report.gauge_band_population,
report.priced_population
);
for row in report.rows.iter() {
println!("[#2673 CONTROL]{row}");
}
assert_eq!(
report.gauge_band_population, 1,
"#2673 control: exactly one direction must sit inside the gauge band"
);
assert_eq!(
report.priced_population, 2,
"#2673 control: the other two directions must be priced"
);
assert_eq!(
report.value_gauge_gradient_resolved, 1,
"#2673 control: the predicate must DETECT a gauge direction the gradient resolves \
(λ={:.3e} inside floor {floor:.3e}, μ={:.3e} above {mu_floor:.3e})",
lambdas[0],
lambdas[0] / vbvs[0]
);
assert_eq!(
report.value_priced_gradient_projected, 1,
"#2673 control: the predicate must DETECT a priced direction the gradient projects \
out (λ={:.3e} above floor {floor:.3e}, μ={:.3e} below {mu_floor:.3e})",
lambdas[1],
lambdas[1] / vbvs[1]
);
}
pub(crate) struct TwoFloorState2673 {
pub a: Array2<f64>,
pub b: Array2<f64>,
pub total_t: usize,
}
fn two_floor_state_2673() -> TwoFloorState2673 {
let n = 24usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.expect("#2673 fixture: the periodic atom must build")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.expect("#2673 fixture: the softmax circle assignment must build");
let mut term = SaeManifoldTerm::new(vec![atom], assignment)
.expect("#2673 fixture: the manifold term must build");
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("#2673 fixture: the arrow system must assemble at this rho");
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_dt, _db, cache) = solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options)
.expect("#2673 fixture: the undamped arrow factorization must succeed");
let a = term
.materialize_exact_hessian_dense(&rho, target.view(), &cache)
.expect("dense exact observed information");
let dim = a.nrows();
let total_t = cache.delta_t_len();
let mut b = Array2::<f64>::zeros((dim, dim));
for col in 0..dim {
let mut vt = Array1::<f64>::zeros(total_t);
let mut vb = Array1::<f64>::zeros(cache.k);
if col < total_t {
vt[col] = 1.0;
} else {
vb[col - total_t] = 1.0;
}
let (out_t, out_b) = gam_solve::arrow_schur::matrix_free_arrow_operator_apply(
&sys,
&cache,
vt.view(),
vb.view(),
)
.expect("majorizer apply");
for row in 0..total_t {
b[[row, col]] = out_t[row];
}
for row in 0..cache.k {
b[[total_t + row, col]] = out_b[row];
}
}
for i in 0..dim {
for j in (i + 1)..dim {
let avg = 0.5 * (b[[i, j]] + b[[j, i]]);
b[[i, j]] = avg;
b[[j, i]] = avg;
}
}
TwoFloorState2673 { a, b, total_t }
}
#[test]
fn the_two_floors_are_incommensurable_thresholds_on_one_operator_2673() {
use gam_linalg::faer_ndarray::strict_symmetric_eigh;
let state = two_floor_state_2673();
let (a_eigs, a_vecs) = strict_symmetric_eigh(&state.a, Side::Lower).expect("A spectrum");
let dim = a_eigs.len();
let max_eig = a_eigs.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let value_threshold = RETIRED_ABSOLUTE_PD_FLOOR_REL_2673 * max_eig.max(1.0);
let mu_floor = f64::EPSILON.sqrt();
let mut ratios = Vec::with_capacity(dim);
let mut gradient_thresholds = Vec::with_capacity(dim);
for index in 0..dim {
let v = a_vecs.column(index);
let vbv = v.dot(&state.b.dot(&v));
assert!(
vbv.is_finite() && vbv > 0.0,
"#2673: B must be positive definite along every direction of A \
(direction {index}: vᵀBv={vbv:.6e})"
);
let gradient_threshold = mu_floor * vbv;
gradient_thresholds.push(gradient_threshold);
ratios.push(gradient_threshold / value_threshold);
}
let min_ratio = ratios.iter().cloned().fold(f64::INFINITY, f64::min);
let max_ratio = ratios.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let min_gradient = gradient_thresholds
.iter()
.cloned()
.fold(f64::INFINITY, f64::min);
let max_gradient = gradient_thresholds
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
println!(
"[#2673 THRESHOLD] dim={dim} λ_max={max_eig:.6e}\n\
[#2673 THRESHOLD] value rule : |λ| ≤ {value_threshold:.6e} (one number, all {dim} directions)\n\
[#2673 THRESHOLD] gradient rule : |λ| < √ε·vᵀBv ∈ [{min_gradient:.6e}, {max_gradient:.6e}]\n\
[#2673 THRESHOLD] ratio gradient/value ∈ [{min_ratio:.6e}, {max_ratio:.6e}] spread={:.6e}x",
max_ratio / min_ratio,
);
assert!(
max_ratio / min_ratio > 1.0 + 1.0e-9,
"#2673: if the two thresholds were one rule in two spellings their ratio would be \
constant across the directions of one operator; measured spread {:.6e}x",
max_ratio / min_ratio
);
assert!(
min_ratio < 1.0 && max_ratio > 1.0,
"#2673: the ratio must straddle 1 for this state to witness that NEITHER rule is \
uniformly the stricter one (ratio ∈ [{min_ratio:.6e}, {max_ratio:.6e}])"
);
}
#[test]
fn the_classification_is_invariant_under_a_reparametrization_2673() {
use gam_linalg::faer_ndarray::{FaerCholesky, strict_symmetric_eigh};
fn min_abs_pencil_curvature(a: &Array2<f64>, b: &Array2<f64>) -> f64 {
let dim = a.nrows();
let lower = b
.cholesky(Side::Lower)
.expect("B is the arrow factorization's own operator and is positive definite")
.lower_triangular();
let forward = |x: &Array2<f64>| -> Array2<f64> {
let mut y = Array2::<f64>::zeros((dim, dim));
for column in 0..dim {
for row in 0..dim {
let mut acc = x[[row, column]];
for k in 0..row {
acc -= lower[[row, k]] * y[[k, column]];
}
y[[row, column]] = acc / lower[[row, row]];
}
}
y
};
let whitened = forward(&forward(a).t().to_owned());
let symmetric = (&whitened + &whitened.t()) * 0.5;
let (mu, _) = strict_symmetric_eigh(&symmetric, Side::Lower).expect("pencil spectrum");
mu.iter().map(|value| value.abs()).fold(f64::INFINITY, f64::min)
}
let state = two_floor_state_2673();
let dim = state.a.nrows();
let (a_eigs, a_vecs) = strict_symmetric_eigh(&state.a, Side::Lower).expect("A spectrum");
let border_width = dim - state.total_t;
assert!(
border_width > 0,
"#2673: the fixture must carry a β border for a border rescaling to be a rescaling"
);
let scale = 1.0e-4_f64;
let mut d = Array1::<f64>::ones(dim);
for index in state.total_t..dim {
d[index] = scale;
}
let congruent = |m: &Array2<f64>| -> Array2<f64> {
Array2::from_shape_fn((dim, dim), |(row, column)| d[row] * m[[row, column]] * d[column])
};
let a_scaled = congruent(&state.a);
let b_scaled = congruent(&state.b);
let mu_plain = min_abs_pencil_curvature(&state.a, &state.b);
let mu_scaled = min_abs_pencil_curvature(&a_scaled, &b_scaled);
let identifiability = f64::EPSILON.sqrt();
println!(
"[#2673 INVARIANCE] border scaled by {scale:.1e} (dim={dim}, border_width={border_width})\n\
[#2673 INVARIANCE] min|μ| plain = {mu_plain:.12e}\n\
[#2673 INVARIANCE] min|μ| scaled = {mu_scaled:.12e} relative move = {:.3e}",
(mu_scaled - mu_plain).abs() / mu_plain
);
assert!(
(mu_scaled - mu_plain).abs() <= 1.0e-6 * mu_plain,
"#2673: the shipped classification is the pencil curvature and MUST be invariant under a \
congruence (plain {mu_plain:.12e}, scaled {mu_scaled:.12e})"
);
assert!(
mu_plain > identifiability && mu_scaled > identifiability,
"#2673: the shipped rule must resolve every direction in BOTH frames for the retired \
rule's flip below to be a disagreement rather than a shared verdict \
(min|μ| plain {mu_plain:.6e}, scaled {mu_scaled:.6e}, floor {identifiability:.6e})"
);
let retired_pinned = |eigs: &Array1<f64>| -> usize {
let max_eig = eigs.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let floor = RETIRED_ABSOLUTE_PD_FLOOR_REL_2673 * max_eig.max(1.0);
eigs.iter().filter(|value| value.abs() <= floor).count()
};
let (scaled_eigs, scaled_vecs) =
strict_symmetric_eigh(&a_scaled, Side::Lower).expect("scaled A spectrum");
let pinned_plain = retired_pinned(&a_eigs);
let pinned_scaled = retired_pinned(&scaled_eigs);
println!(
"[#2673 INVARIANCE] retired rule pins {pinned_plain} of {dim} directions before the \
rescaling and {pinned_scaled} after"
);
assert_eq!(
pinned_plain, 0,
"#2673: the retired rule pins nothing on the unscaled state, which is why every previous \
crossing count on it was structurally zero"
);
assert!(
pinned_scaled > 0,
"#2673: rescaling the border must push at least one direction under the retired absolute \
band, or this state does not witness the crossing"
);
let shipped_pinned = |eigs: &Array1<f64>,
vecs: &Array2<f64>,
b: &Array2<f64>,
label: &str|
-> (usize, usize) {
let norm = eigs.iter().map(|value| value.abs()).fold(0.0_f64, f64::max);
let arithmetic = (dim as f64) * f64::EPSILON * norm;
let mut pinned = 0usize;
let mut arithmetic_binds = 0usize;
let mut worst_margin = f64::INFINITY;
for index in 0..dim {
let v = vecs.column(index);
let vbv = v.dot(&b.dot(&v));
let floor = arithmetic.max(identifiability * vbv);
if eigs[index].abs() <= floor {
pinned += 1;
}
if arithmetic > identifiability * vbv {
arithmetic_binds += 1;
worst_margin = worst_margin.min(eigs[index].abs() / arithmetic);
}
}
println!(
"[#2673 INVARIANCE] shipped rule on the {label} frame: pins {pinned} of {dim}; the \
arithmetic term binds on {arithmetic_binds} (there |λ| still clears it by \
{worst_margin:.3e}x, so the verdict is the identifiability verdict)"
);
(pinned, arithmetic_binds)
};
let (plain_pinned, _) = shipped_pinned(&a_eigs, &a_vecs, &state.b, "plain");
let (scaled_pinned, scaled_arithmetic_binds) =
shipped_pinned(&scaled_eigs, &scaled_vecs, &b_scaled, "rescaled");
assert_eq!(
(plain_pinned, scaled_pinned),
(0, 0),
"#2673: the shipped rule must reach the SAME verdict on every direction in both frames; \
a units change is not a modelling decision"
);
let retired_floor = RETIRED_ABSOLUTE_PD_FLOOR_REL_2673
* scaled_eigs
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max)
.max(1.0);
let mut crossings = 0usize;
for index in 0..dim {
let v = scaled_vecs.column(index);
let vbv = v.dot(&b_scaled.dot(&v));
let mu = scaled_eigs[index] / vbv;
if scaled_eigs[index].abs() <= retired_floor && mu.abs() >= identifiability {
crossings += 1;
println!(
"[#2673 INVARIANCE] direction {index}: RETIRED VALUE RULE=gauge \
(|λ|={:.6e} ≤ {retired_floor:.6e}) but GRADIENT=resolved \
(|μ|={:.6e} ≥ {identifiability:.6e})",
scaled_eigs[index].abs(),
mu.abs()
);
}
}
assert!(
crossings > 0,
"#2673: the rescaled frame must exhibit the `value=gauge & gradient=resolved` crossing \
the issue was filed about"
);
assert!(
scaled_arithmetic_binds > 0,
"#2673: this frame is chosen so the arithmetic term BINDS somewhere, or the paragraph \
above is untested commentary"
);
}
#[test]
fn two_floors_overlap_region_direction_count_2673() {
use gam_linalg::faer_ndarray::strict_symmetric_eigh;
let n = 24usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_dt, _db, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let a = term
.materialize_exact_hessian_dense(&rho, target.view(), &cache)
.expect("dense exact observed information");
let dim = a.nrows();
let total_t = cache.delta_t_len();
let mut b = Array2::<f64>::zeros((dim, dim));
for col in 0..dim {
let mut vt = Array1::<f64>::zeros(total_t);
let mut vb = Array1::<f64>::zeros(cache.k);
if col < total_t {
vt[col] = 1.0;
} else {
vb[col - total_t] = 1.0;
}
let (out_t, out_b) =
gam_solve::arrow_schur::matrix_free_arrow_operator_apply(&sys, &cache, vt.view(), vb.view())
.expect("majorizer apply");
for row in 0..total_t {
b[[row, col]] = out_t[row];
}
for row in 0..cache.k {
b[[total_t + row, col]] = out_b[row];
}
}
for i in 0..dim {
for j in (i + 1)..dim {
let avg = 0.5 * (b[[i, j]] + b[[j, i]]);
b[[i, j]] = avg;
b[[j, i]] = avg;
}
}
let (a_eigs, a_vecs) = strict_symmetric_eigh(&a, Side::Lower).expect("A spectrum");
let max_eig = a_eigs.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let floor = RETIRED_ABSOLUTE_PD_FLOOR_REL_2673 * max_eig.max(1.0);
let mu_floor = f64::EPSILON.sqrt();
let report = classify_two_floor_overlap_2673(a_eigs.view(), a_vecs.view(), &b, floor, mu_floor);
let min_eig = a_eigs.iter().cloned().fold(f64::INFINITY, f64::min);
println!(
"[#2673 PENCIL] dim={dim} λ_max={max_eig:.6e} λ_min={min_eig:.6e} \
floor={floor:.6e} μ_floor={mu_floor:.6e}\n\
[#2673 PENCIL] value=gauge & gradient=resolved : {}\n\
[#2673 PENCIL] value=priced & gradient=projected: {}\n\
[#2673 PENCIL] directions that change class TOTAL: {}",
report.value_gauge_gradient_resolved,
report.value_priced_gradient_projected,
report.value_gauge_gradient_resolved + report.value_priced_gradient_projected,
);
println!(
"[#2673 RISK] gauge-band population (|λ| ≤ floor) : {} of {dim}\n\
[#2673 RISK] priced population (|λ| > floor) : {} of {dim}\n\
[#2673 RISK] |λ| range : [{:.6e}, {:.6e}] (floor={floor:.6e}, min|λ|/floor={:.6e})\n\
[#2673 RISK] vᵀBv range : [{:.6e}, {:.6e}]\n\
[#2673 RISK] a gauge direction crosses iff vᵀBv ≤ |λ|/μ_floor ≤ {:.6e}\n\
[#2673 RISK] a priced direction crosses iff vᵀBv ≥ |λ|/μ_floor ≥ {:.6e}",
report.gauge_band_population,
report.priced_population,
report.min_abs_lambda,
report.max_abs_lambda,
report.min_abs_lambda / floor,
report.min_vbv,
report.max_vbv,
floor / mu_floor,
report.min_abs_lambda_priced / mu_floor,
);
for row in report.rows.iter().take(20) {
println!("[#2673 PENCIL]{row}");
}
if report.gauge_band_population == 0 {
println!(
"[#2673 VACUOUS] no direction lies in the gauge band, so \
`value=gauge & gradient=resolved` = {} is STRUCTURALLY zero here and is \
NOT evidence that the two rules agree",
report.value_gauge_gradient_resolved
);
}
if report.max_vbv < report.min_abs_lambda_priced / mu_floor {
println!(
"[#2673 VACUOUS] max vᵀBv = {:.6e} is below the {:.6e} any priced direction \
would need, so `value=priced & gradient=projected` = {} is STRUCTURALLY zero \
here and is NOT evidence that the two rules agree",
report.max_vbv,
report.min_abs_lambda_priced / mu_floor,
report.value_priced_gradient_projected
);
}
assert_eq!(a_eigs.len(), dim, "#2673: the A spectrum must span every direction");
assert!(
a_eigs.iter().all(|v| v.is_finite()),
"#2673: every A eigenvalue must be finite, or the classification is undefined"
);
assert!(
floor > 0.0 && floor.is_finite(),
"#2673: the absolute floor must be a real positive band, got {floor}"
);
assert_eq!(
report.gauge_band_population + report.priced_population,
dim,
"#2673: the gauge band and its complement must partition the {dim} directions"
);
assert!(
min_eig < -floor,
"#2673: the fixture must be decisively indefinite to be the adversarial \
population this issue is about (λ_min={min_eig:.6e}, floor={floor:.6e})"
);
}
#[test]
fn exact_a_ard_operator_derivative_is_the_unmajorized_hessian_2515() {
use crate::manifold::arrow_solver::SaeLocalRowVar;
let n = 24usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords.clone()],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_dt, _db, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let alpha = term.validated_ard_precisions(&rho).unwrap()[0][0];
let period = Some(1.0_f64);
let mut any_clamped = 0usize;
for row in 0..n {
let t = coords[[row, 0]];
let prior = ArdAxisPrior::eval(alpha, t, period);
assert_eq!(
prior.hess,
prior.psd_majorizer_hess() + prior.negative_hessian_remainder(),
"#2515: the exact/majorizer split must hold BIT-FOR-BIT at row {row}; the \
one-token substitution psd_majorizer_hess() -> hess is only the exact-A \
operator derivative because of this identity"
);
if prior.negative_hessian_remainder() != 0.0 {
any_clamped += 1;
}
}
assert!(
any_clamped > 0,
"#2515: this fixture must have rows where the periodic ARD clamp is ACTIVE \
(cos κt < 0), or the exact and majorized curvatures are equal and this test \
proves nothing about the substitution"
);
let delta_by_flat = term
.exact_stationarity_penalty_derivative_delta_by_flat(&rho, &cache)
.expect("the exact-A penalty derivative delta must be assemblable");
let ard_flat = rho.ard_flat_index(0, 0);
let delta = delta_by_flat
.get(&ard_flat)
.expect("#2515: the ARD coordinate must carry a ∂ΔC/∂ρ block on this fixture");
let row_weights = term.row_loss_weights.clone();
let mut checked = 0usize;
for row in 0..n {
let base = cache.row_offsets[row];
let vars = term.row_vars_for_cache_row(row, &cache).unwrap();
for (local, var) in vars.iter().enumerate() {
let SaeLocalRowVar::Coord { atom, axis } = *var else {
continue;
};
if atom != 0 || axis != 0 {
continue;
}
let w_row = row_weights.as_ref().map_or(1.0, |w| w[row]);
let expected =
w_row * ArdAxisPrior::eval(alpha, coords[[row, 0]], period).negative_hessian_remainder();
let actual = delta[[base + local, base + local]];
assert!(
(actual - expected).abs() <= 1.0e-12 * expected.abs().max(1.0),
"#2515: the production ∂ΔC/∂ρ_ard map must equal \
w_row·negative_hessian_remainder at row {row} (map={actual:.17e}, \
closed form={expected:.17e}). These are two independent derivations of \
the same quantity; if they disagree the substitution is unsound."
);
checked += 1;
}
}
assert_eq!(
checked, n,
"#2515: every row must contribute an ARD coordinate slot to the check"
);
println!(
"[#2515 B-FULL step 1] alpha={alpha:.6e} rows={n} rows with an ACTIVE clamp={any_clamped}\n\
[#2515 B-FULL step 1] verified on all {checked} rows: \
d(A)/d(rho_ard) = psd_majorizer_hess + negative_hessian_remainder = prior.hess"
);
}
#[test]
fn from_probes_exact_a_theta_adjoint_matches_dense_2515() {
use crate::manifold::construction::ThetaAdjointDhChannel;
let n = 24usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_dt, _db, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let k = cache.k;
let sqrt_k = (k as f64).sqrt();
let probes: Vec<Array1<f64>> = (0..k)
.map(|j| {
let mut v = Array1::<f64>::zeros(k);
v[j] = sqrt_k;
v
})
.collect();
let sinv: Vec<Array1<f64>> = probes
.iter()
.map(|v| cache.schur_inverse_apply(v.view()).unwrap())
.collect();
let solver = DeflatedArrowSolver::plain(&cache);
let inv = term.materialize_joint_inverse(&cache, &solver).unwrap();
let dense_exact = term
.logdet_theta_adjoint_dense(
&rho,
&cache,
&inv,
ThetaAdjointDhChannel::All,
true,
true,
None,
)
.expect("dense exact-A theta adjoint");
let dense_majorizer = term
.logdet_theta_adjoint_dense(
&rho,
&cache,
&inv,
ThetaAdjointDhChannel::All,
true,
false,
None,
)
.expect("dense majorizer-B theta adjoint");
let separation = dense_exact
.t
.iter()
.zip(dense_majorizer.t.iter())
.chain(dense_exact.beta.iter().zip(dense_majorizer.beta.iter()))
.map(|(a, b)| (a - b).abs())
.fold(0.0_f64, f64::max);
assert!(
separation > 1.0e-6,
"#2515: the exact-A and majorizer-B theta adjoints must SEPARATE on this \
fixture (max|Δ|={separation:.6e}); without separation this test cannot \
distinguish a correct exact-A port from one that ignores the flag"
);
let probes_exact = term
.logdet_theta_adjoint_from_probes(
&rho,
&cache,
&probes,
&sinv,
EvidenceOperator::ExactObservedInformation,
None,
)
.expect("from-probes exact-A theta adjoint");
let probes_majorizer = term
.logdet_theta_adjoint_from_probes(
&rho,
&cache,
&probes,
&sinv,
EvidenceOperator::Majorizer,
None,
)
.expect("from-probes majorizer-B theta adjoint");
let worst = |x: &SaeArrowVector, y: &SaeArrowVector| -> f64 {
x.t.iter()
.zip(y.t.iter())
.chain(x.beta.iter().zip(y.beta.iter()))
.map(|(a, b)| (a - b).abs())
.fold(0.0_f64, f64::max)
};
let exact_gap = worst(&probes_exact, &dense_exact);
let majorizer_gap = worst(&probes_majorizer, &dense_majorizer);
println!(
"[#2515 B-FULL theta] exact-A |from_probes - dense| = {exact_gap:.6e} \
majorizer-B |from_probes - dense| = {majorizer_gap:.6e} \
exact-vs-majorizer separation = {separation:.6e}"
);
assert!(
majorizer_gap <= 1.0e-9,
"#2515: the majorizer-B from-probes adjoint must still match its dense \
reference (max|Δ|={majorizer_gap:.6e}); the exact-A port must not disturb it"
);
assert!(
exact_gap <= 1.0e-9,
"#2515: the from-probes exact-A theta adjoint must reproduce the dense \
exact-A adjoint term for term (max|Δ|={exact_gap:.6e}, separation from the \
majorizer is {separation:.6e}). A mirrored leg with the wrong sign, jet, or \
index pair lands here."
);
}
#[test]
fn zz_measure_smoothness_dof_bundle_vs_deflated_2499() {
let n = 24usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_dt, _db, cache) = solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let k = cache.k;
let lambda_smooth = rho.lambda_smooth_vec().unwrap();
let plain = cache.schur_inverse_block(0..k).unwrap();
let deflated = cache.schur_inverse_block_deflated(0..k).unwrap();
let mut max_abs = 0.0_f64;
for i in 0..k {
for j in 0..k {
max_abs = max_abs.max((plain[[i, j]] - deflated[[i, j]]).abs());
}
}
let sqrt_k = (k as f64).sqrt();
let probes: Vec<Array1<f64>> = (0..k)
.map(|j| {
let mut v = Array1::<f64>::zeros(k);
v[j] = sqrt_k;
v
})
.collect();
let sinv: Vec<Array1<f64>> = probes
.iter()
.map(|v| cache.schur_inverse_apply(v.view()).unwrap())
.collect();
let from_probes = term
.decoder_smoothness_effective_dof_per_atom_from_probes(&probes, &sinv, &lambda_smooth)
.unwrap();
let dense_deflated = term
.decoder_smoothness_effective_dof_per_atom(&cache, &lambda_smooth)
.unwrap();
let solver = DeflatedArrowSolver::plain(&cache);
let with_solver = term
.decoder_smoothness_effective_dof_with_solver_per_atom(&cache, &solver, &lambda_smooth)
.unwrap();
println!(
"[#2499 smoothness-dof] k={k} max|S^-1_plain - S^-1_deflated|={max_abs:.6e}\n\
from_probes={from_probes:?}\n dense_deflated={dense_deflated:?}\n with_solver={with_solver:?}\n\
plain_selected_inverse_available={}",
solver.plain_selected_inverse_available()
);
assert!(
max_abs.is_finite() && max_abs >= 0.0,
"[#2499] the plain-vs-deflated selected-inverse gap is a max |·| and must be a finite \
non-negative magnitude, got {max_abs}"
);
assert!(
!from_probes.is_empty()
&& from_probes.len() == dense_deflated.len()
&& from_probes.len() == with_solver.len(),
"[#2499] the three smoothness-EDF routes must return the same nonempty per-atom shape: \
from_probes={from_probes:?} dense_deflated={dense_deflated:?} with_solver={with_solver:?}"
);
assert!(
from_probes
.iter()
.chain(dense_deflated.iter())
.chain(with_solver.iter())
.all(|v| v.is_finite() && *v >= 0.0),
"[#2499] a smoothness effective dof is a finite non-negative trace: \
from_probes={from_probes:?} dense_deflated={dense_deflated:?} with_solver={with_solver:?}"
);
let loss = term.loss(target.view(), &rho).unwrap();
let after_loss = term
.decoder_smoothness_effective_dof_with_solver_per_atom(&cache, &solver, &lambda_smooth)
.unwrap();
let dense_components = term
.analytic_outer_rho_gradient_components(target.view(), &rho, &loss, &cache, &solver)
.unwrap();
let bundled_components = term
.analytic_outer_rho_gradient_components_with_bundle(
target.view(),
&rho,
&loss,
&cache,
&solver,
Some(BundleEvidenceGeometry {
operator: EvidenceOperator::Majorizer,
cache: &cache,
probes: &probes,
sinv: &sinv,
}),
None,
)
.unwrap();
let smooth_index = rho.smooth_flat_index(0);
println!(
"[#2499 smoothness-dof] after term.loss(): with_solver={after_loss:?}\n\
logdet_trace[smooth {smooth_index}] dense={:.17e} bundled={:.17e}\n\
half-of-with_solver={:.17e} half-of-from_probes={:.17e}",
dense_components.logdet_trace[smooth_index],
bundled_components.logdet_trace[smooth_index],
0.5 * after_loss[0],
0.5 * from_probes[0],
);
assert!(
smooth_index < dense_components.logdet_trace.len()
&& smooth_index < bundled_components.logdet_trace.len(),
"[#2499] the smooth coordinate {smooth_index} must be present in both logdet-trace \
vectors (dense len={}, bundled len={})",
dense_components.logdet_trace.len(),
bundled_components.logdet_trace.len()
);
assert!(
dense_components.logdet_trace[smooth_index].is_finite()
&& bundled_components.logdet_trace[smooth_index].is_finite()
&& after_loss.iter().all(|v| v.is_finite()),
"[#2499] both logdet-trace routes must report a finite smooth coordinate \
(dense={}, bundled={}, after_loss={after_loss:?})",
dense_components.logdet_trace[smooth_index],
bundled_components.logdet_trace[smooth_index]
);
}
#[test]
fn solve_exact_stationarity_is_self_adjoint_2080() {
let n = 24usize;
let p = 2usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let decoder = array![[0.30, -0.10], [1.20, 0.20], [0.10, 1.10]];
assert_eq!(decoder.ncols(), p);
let mut target = phi.dot(&decoder);
for row in 0..n {
target[[row, 0]] += 1.0e-3 * (0.37 * row as f64).sin();
target[[row, 1]] += 1.0e-3 * (0.29 * row as f64).cos();
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_delta_t, _delta_beta, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let n_params = rho.to_flat().len();
assert!(n_params >= 2, "fixture must expose ≥2 outer coordinates");
let mut u = term.outer_rho_gradient_ift_rhs(&rho, 0, &cache).unwrap();
let mut v = term
.outer_rho_gradient_ift_rhs(&rho, n_params - 1, &cache)
.unwrap();
for index in 0..u.t.len() {
u.t[index] += ((index as f64 + 1.0) * 0.37).sin();
v.t[index] += ((index as f64 + 2.0) * 0.29).cos();
}
for index in 0..u.beta.len() {
u.beta[index] += ((index as f64 + 3.0) * 0.23).cos();
v.beta[index] += ((index as f64 + 4.0) * 0.31).sin();
}
let a_u = term
.solve_exact_stationarity(&rho, target.view(), &cache, &u)
.unwrap();
let a_v = term
.solve_exact_stationarity(&rho, target.view(), &cache, &v)
.unwrap();
let lhs = a_u
.t
.iter()
.zip(v.t.iter())
.map(|(a, b)| a * b)
.sum::<f64>()
+ a_u
.beta
.iter()
.zip(v.beta.iter())
.map(|(a, b)| a * b)
.sum::<f64>();
let rhs =
u.t.iter()
.zip(a_v.t.iter())
.map(|(a, b)| a * b)
.sum::<f64>()
+ u.beta
.iter()
.zip(a_v.beta.iter())
.map(|(a, b)| a * b)
.sum::<f64>();
let scale = lhs.abs().max(rhs.abs()).max(1.0);
assert!(
(lhs - rhs).abs() <= 1.0e-6 * scale,
"solve_exact_stationarity must be self-adjoint (the #2080(A) single-adjoint \
IFT identity): ⟨A⁺u,v⟩={lhs} vs ⟨u,A⁺v⟩={rhs}"
);
let response_norm_sq = a_u.t.dot(&a_u.t)
+ a_u.beta.dot(&a_u.beta)
+ a_v.t.dot(&a_v.t)
+ a_v.beta.dot(&a_v.beta);
assert!(
response_norm_sq.is_finite() && response_norm_sq.sqrt() > f64::EPSILON.sqrt(),
"self-adjoint pin must have a non-trivial finite quotient response: \
lhs={lhs} rhs={rhs} response_norm_sq={response_norm_sq}"
);
}
#[test]
pub(crate) fn latent_block_inverse_diagonal_hutchinson_matches_exact_trace() {
let n = 24usize;
let p = 2usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.25) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.30, -0.10], [0.20, 0.40], [-0.35, 0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = Array2::from_shape_fn((n, p), |(row, col)| {
let x = (row as f64 + 0.5) / n as f64;
if col == 0 {
0.45 * (std::f64::consts::TAU * x).sin() + 0.07
} else {
-0.20 * (std::f64::consts::TAU * x).cos() + 0.03 * row as f64
}
});
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![250.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let (_delta_t, _delta_beta, cache) =
solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let exact = cache.latent_block_inverse_diagonal().unwrap();
let hutch =
SaeManifoldTerm::latent_block_inverse_diagonal_hutchinson(&cache, 20_000, 0xABCD_1234)
.unwrap();
assert_eq!(exact.len(), hutch.len());
let exact_trace: f64 = exact.iter().sum();
let hutch_trace: f64 = hutch.iter().sum();
assert!(
(hutch_trace - exact_trace).abs() <= 0.02 * exact_trace.abs().max(1.0e-6),
"Hutchinson latent trace {hutch_trace} vs exact {exact_trace} exceeds 2% tol"
);
let coord_offsets = term.assignment.coord_offsets();
let block_start = coord_offsets[0];
let mut exact_axis0 = 0.0_f64;
let mut hutch_axis0 = 0.0_f64;
match term.last_row_layout {
Some(ref layout) => {
for row in 0..n {
let row_base = cache.row_offsets[row];
if let Some(pos) = layout.active_atoms[row].iter().position(|&k| k == 0) {
let s = row_base + layout.coord_starts[row][pos];
exact_axis0 += exact[s];
hutch_axis0 += hutch[s];
}
}
}
None => {
for row in 0..n {
let s = cache.row_offsets[row] + block_start;
exact_axis0 += exact[s];
hutch_axis0 += hutch[s];
}
}
}
assert!(
(hutch_axis0 - exact_axis0).abs() <= 0.05 * exact_axis0.abs().max(1.0e-6),
"Hutchinson ARD axis trace {hutch_axis0} vs exact {exact_axis0} exceeds 5% tol"
);
assert!(
exact_axis0 > 0.0,
"posterior-variance trace must be positive (sanity on the fixture)"
);
}
#[test]
pub(crate) fn streaming_plan_routes_by_memory_budget_with_identical_logdet() {
let (term0, target, rho) = small_two_atom_periodic_term();
let total_basis: usize = term0.atoms.iter().map(|atom| atom.basis_size()).sum();
let d_max = term0
.atoms
.iter()
.map(SaeManifoldAtom::latent_dim)
.max()
.unwrap();
let dense_plan = sae_streaming_plan_from_budget(
term0.n_obs(),
total_basis,
term0.k_atoms(),
d_max,
term0.beta_dim(),
usize::MAX / 4,
1024 * 1024,
usize::MAX / 2,
);
assert!(!dense_plan.streaming);
assert!(dense_plan.direct_admitted);
let streaming_plan = sae_streaming_plan_from_budget(
term0.n_obs(),
total_basis,
term0.k_atoms(),
d_max,
term0.beta_dim(),
1,
512,
2,
);
assert!(streaming_plan.streaming);
assert!(!streaming_plan.direct_admitted);
let mut full = term0.clone();
full.guards_enabled = false;
full.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
2,
0.25,
1.0e-4,
1.0e-4,
)
.unwrap();
let sys = full
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let options = ArrowSolveOptions::direct().with_positive_definite_evidence();
let factor_result = solve_arrow_newton_step_with_options(&sys, 0.0, 0.0, &options).unwrap();
let full_logdet = arrow_log_det_from_cache(&factor_result.2).unwrap();
let mut streaming = StreamingArrowSchur::from_system(&sys, streaming_plan.chunk_size);
let streaming_logdet = streaming.exact_arrow_log_det(0.0, 0.0, &options).unwrap();
assert_abs_diff_eq!(streaming_logdet, full_logdet, epsilon = 1.0e-8);
}
#[test]
pub(crate) fn giant_host_working_set_plan_flips_to_matrix_free_before_dense_allocation() {
let n_obs = 128usize;
let total_basis = 48usize;
let k_atoms = 8usize;
let d_max = 2usize;
let p_out = 2048usize;
let border_dim = total_basis * p_out;
let budget = 60usize * 1024 * 1024 * 1024;
let plan = sae_streaming_plan_from_budget(
n_obs,
total_basis,
k_atoms,
d_max,
border_dim,
budget,
SAE_CPU_L2_CACHE_BYTES * SAE_CHUNK_CACHE_MULTIPLE,
120usize * 1024 * 1024 * 1024,
);
assert_eq!(border_dim, 98_304);
assert_eq!(
plan.estimated_row_cross_bytes,
n_obs * k_atoms * (1 + d_max) * border_dim * SAE_BYTES_PER_F64
);
assert!(plan.estimated_dense_schur_bytes > budget);
assert!(plan.estimated_matrix_free_peak_bytes < budget);
assert!(plan.streaming);
assert!(!plan.direct_admitted);
assert!(plan.matrix_free_admitted);
assert_eq!(
plan.solve_options_for_border_dim(border_dim).mode,
gam_solve::arrow_schur::ArrowSolverMode::InexactPCG
);
}
#[test]
pub(crate) fn matrix_free_plan_refuses_genuinely_exhausted_process_budget() {
let n_obs = 508usize;
let total_basis = 6usize;
let k_atoms = 1usize;
let d_max = 1usize;
let border_dim = 32usize;
let plan = sae_streaming_plan_from_budget(
n_obs,
total_basis,
k_atoms,
d_max,
border_dim,
0, SAE_CPU_L2_CACHE_BYTES * SAE_CHUNK_CACHE_MULTIPLE,
0, );
assert!(!plan.direct_admitted);
assert!(plan.streaming);
assert!(!plan.matrix_free_admitted);
assert!(plan.admitted_or_error(n_obs, border_dim, k_atoms).is_err());
}
fn planted_topk_sae_term(
n: usize,
k_atoms: usize,
planted: &[Vec<usize>],
p: usize,
) -> (SaeManifoldTerm, Array2<f64>) {
assert_eq!(planted.len(), n);
let mut atoms = Vec::with_capacity(k_atoms);
let mut coord_blocks = Vec::with_capacity(k_atoms);
let mut manifolds = Vec::with_capacity(k_atoms);
let coords = Array2::<f64>::from_shape_fn((n, 1), |(row, _)| (row as f64 / n as f64) - 0.5);
for atom_idx in 0..k_atoms {
let mut phi = Array2::<f64>::zeros((n, 2));
let mut jet = Array3::<f64>::zeros((n, 2, 1));
for row in 0..n {
phi[[row, 0]] = 1.0;
phi[[row, 1]] = coords[[row, 0]];
jet[[row, 1, 0]] = 1.0;
}
let mut decoder = Array2::<f64>::zeros((2, p));
decoder[[1, atom_idx % p]] = 0.1 + 0.01 * ((atom_idx % 7) as f64);
atoms.push(
SaeManifoldAtom::new_with_provided_function_gram(
format!("atom{atom_idx}"),
SaeAtomBasisKind::EuclideanPatch,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(2),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction"),
);
coord_blocks.push(coords.clone());
manifolds.push(LatentManifold::Euclidean);
}
let mut logits = Array2::<f64>::from_elem((n, k_atoms), -6.0);
for (row, active) in planted.iter().enumerate() {
for &k in active {
logits[[row, k]] = 6.0;
}
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coord_blocks,
manifolds,
AssignmentMode::top_k_support(planted[0].len()),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
let term = SaeManifoldTerm::new(atoms, assignment)
.expect("term fixture: every atom's row count matches the assignment's");
let target = Array2::<f64>::from_shape_fn((n, p), |(row, c)| 0.05 * ((row + c) as f64).sin());
(term, target)
}
#[test]
pub(crate) fn large_k_topk_encode_is_support_bounded_and_exact() {
let n = 8usize;
let p = 4usize;
let top_k = 3usize;
let planted: Vec<Vec<usize>> = (0..n).map(|row| vec![row, 300 + row, 700 + row]).collect();
let assemble_dims = |k_atoms: usize| -> (Vec<usize>, Vec<Vec<usize>>) {
let (mut term, target) = planted_topk_sae_term(n, k_atoms, &planted, p);
term.fixed_decoder_assembly = true;
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1); k_atoms]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("exact TopK fixed-decoder assembly must succeed at large K");
let dims: Vec<usize> = sys.rows.iter().map(|r| r.htt.nrows()).collect();
for r in &sys.rows {
assert_eq!(r.htt.nrows(), r.htt.ncols());
assert_eq!(r.htt.nrows(), r.gt.len());
}
let layout = term
.last_row_layout
.clone()
.expect("TopK must install its exact compact support layout");
let active: Vec<Vec<usize>> = layout.active_atoms.clone();
(dims, active)
};
let (dims_1k, active_1k) = assemble_dims(1_000);
let (dims_10k, active_10k) = assemble_dims(10_000);
let bound = top_k; for row in 0..n {
assert!(
dims_1k[row] <= bound,
"row {row} K=1000 compact dim {} exceeds O(top_k) bound {bound}",
dims_1k[row]
);
assert_eq!(
dims_1k[row], dims_10k[row],
"row {row} compact dim must be INDEPENDENT of total K (n-free contract): \
K=1000 gave {} but K=10000 gave {}",
dims_1k[row], dims_10k[row]
);
}
let compact_work: usize = dims_1k.iter().map(|&q| q * q).sum();
let dense_q = 1_000; let dense_work = n * dense_q * dense_q;
assert!(
compact_work * 100 < dense_work,
"compact work {compact_work} must be << dense work {dense_work}"
);
for row in 0..n {
let mut expected = planted[row].clone();
expected.sort_unstable();
assert_eq!(
active_1k[row], expected,
"row {row} K=1000 active set must recover the planted top-{top_k} support"
);
assert_eq!(
active_10k[row], expected,
"row {row} K=10000 active set must recover the planted top-{top_k} support"
);
}
}
#[test]
pub(crate) fn sparse_active_layout_work_scales_with_active_atoms_not_total_k() {
let n = 3;
let k_atoms = 100_000;
let mut gates = Vec::with_capacity(n);
for row in 0..n {
let mut row_gates = Array1::<f64>::zeros(k_atoms);
for atom in [row, 10_000 + row, 90_000 + row] {
row_gates[atom] = 1.0;
}
gates.push(row_gates);
}
let coord_dims = vec![1usize; k_atoms];
let coord_offsets_full: Vec<usize> = (0..k_atoms).collect();
let layout = SaeRowLayout::from_topk_gates(&gates, 3, coord_dims, coord_offsets_full).unwrap();
for row in 0..n {
assert_eq!(layout.active_atoms[row].len(), 3);
assert_eq!(layout.row_q_active(row), 3);
}
let compact_work: usize = (0..n)
.map(|row| {
let q = layout.row_q_active(row);
q * q
})
.sum();
let dense_q = k_atoms;
let dense_work = n * dense_q * dense_q;
assert!(compact_work < dense_work / 1_000_000_000);
assert_eq!(compact_work, n * 9);
}
#[test]
pub(crate) fn run_joint_fit_arrow_schur_escalates_ridge_on_non_pd_row_block() {
let coords = array![[0.1], [0.4], [0.7]];
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.05], [-0.05], [0.05]],
Array2::<f64>::zeros((3, 3)),
)
.unwrap();
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((3, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.20], [-0.10], [0.45]];
let mut rho = SaeManifoldRho::new(0.0, -20.0, vec![Array1::<f64>::zeros(1)]);
let result =
term.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 1, 1.0, 1.0e-6, 1.0e-6);
assert!(
result.is_ok(),
"run_joint_fit_arrow_schur should recover from degenerate H_tt via LM ridge escalation; got: {result:?}",
);
}
#[test]
pub(crate) fn rank_revealing_reduction_collapses_unexcited_circle_harmonic_to_full_rank() {
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(5).unwrap());
let coords = array![[0.1], [0.45], [0.8], [0.1], [0.45], [0.8]];
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
assert_eq!(
phi.ncols(),
5,
"fixed-depth circle basis emits M = 5 columns"
);
let penalty = Array2::<f64>::eye(5);
let decoder = array![[0.05], [-0.05], [0.05], [0.02], [-0.02]];
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi.clone(),
jet,
decoder.clone(),
penalty,
)
.unwrap()
.with_basis_second_jet(evaluator.clone());
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::ones((6, 1)),
vec![coords.clone()],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let recon_before = phi.dot(&decoder);
term.reduce_atoms_to_data_supported_rank().unwrap();
let r = term.atoms[0].basis_size();
assert_eq!(
r, 3,
"rank-revealing reduction must drop the unexcited harmonic (r = 3 < M = 5)",
);
assert_eq!(term.atoms[0].decoder_coefficients().nrows(), 3);
assert_eq!(term.atoms[0].basis_jacobian.dim(), (6, 3, 1));
assert_eq!(term.atoms[0].smooth_penalty().dim(), (3, 3));
use gam_linalg::faer_ndarray::FaerEigh;
let reduced_design = term.atoms[0].basis_values.clone();
let gram = reduced_design.t().dot(&reduced_design);
let (evals, _) = gram.eigh(faer::Side::Lower).unwrap();
let max_eig = evals.iter().cloned().fold(0.0_f64, f64::max);
for &lam in evals.iter() {
assert!(
lam > 1e-9 * max_eig,
"reduced design Gram must be full rank; got eigenvalue {lam} (max {max_eig})",
);
}
let recon_after = reduced_design.dot(term.atoms[0].decoder_coefficients());
for i in 0..recon_before.nrows() {
assert!(
(recon_before[[i, 0]] - recon_after[[i, 0]]).abs() < 1e-9,
"reduction must not change the data-fit reconstruction at row {i}: \
before={} after={}",
recon_before[[i, 0]],
recon_after[[i, 0]],
);
}
let (refreshed, _) = term.atoms[0]
.basis_evaluator
.as_ref()
.unwrap()
.evaluate(coords.view())
.unwrap();
assert_eq!(
refreshed.ncols(),
3,
"the SubspaceReducedEvaluator must re-emit the reduced width on refresh",
);
for i in 0..refreshed.nrows() {
for j in 0..3 {
assert!(
(refreshed[[i, j]] - reduced_design[[i, j]]).abs() < 1e-12,
"refresh must reproduce the reduced design bit-for-bit",
);
}
}
}
#[test]
pub(crate) fn subspace_reduced_evaluator_composes_all_jets_by_q() {
let inner = Arc::new(PeriodicHarmonicEvaluator::new(7).unwrap());
let coords = array![[-0.3_f64], [0.0], [0.15], [0.42], [0.88]];
let m = inner.num_basis; let mut a = Array2::<f64>::zeros((m, m));
for i in 0..m {
for j in 0..m {
a[[i, j]] = 1.0 / (1.0 + (i as f64 - j as f64).abs());
}
}
let (_evals, evecs) = a.eigh(Side::Lower).unwrap();
let r = 4usize;
let mut q = Array2::<f64>::zeros((m, r));
for col in 0..r {
for row in 0..m {
q[[row, col]] = evecs[[row, col]];
}
}
let reduced = SubspaceReducedEvaluator::new(inner.clone(), q.clone()).unwrap();
assert_eq!(reduced.inner_width(), m);
assert_eq!(reduced.reduced_width(), r);
let (phi_in, jet_in) = inner.evaluate(coords.view()).unwrap();
let (phi_red, jet_red) = reduced.evaluate(coords.view()).unwrap();
let phi_expect = phi_in.dot(&q);
assert_eq!(phi_red.dim(), phi_expect.dim());
for i in 0..phi_red.nrows() {
for j in 0..r {
assert_abs_diff_eq!(phi_red[[i, j]], phi_expect[[i, j]], epsilon = 1e-12);
}
}
for axis in 0..jet_in.shape()[2] {
let expect = jet_in.slice(s![.., .., axis]).to_owned().dot(&q);
for i in 0..jet_red.shape()[0] {
for j in 0..r {
assert_abs_diff_eq!(jet_red[[i, j, axis]], expect[[i, j]], epsilon = 1e-12);
}
}
}
let h_in = inner.second_jet(coords.view()).unwrap();
let h_red = reduced.second_jet(coords.view()).unwrap();
let d = h_in.shape()[2];
for a_ax in 0..d {
for c_ax in 0..d {
let expect = h_in.slice(s![.., .., a_ax, c_ax]).to_owned().dot(&q);
for i in 0..h_red.shape()[0] {
for j in 0..r {
assert_abs_diff_eq!(h_red[[i, j, a_ax, c_ax]], expect[[i, j]], epsilon = 1e-12);
}
}
}
}
let t_in = inner.third_jet(coords.view()).unwrap();
let t_red = reduced.third_jet_dyn(coords.view()).unwrap().unwrap();
for a_ax in 0..d {
for c_ax in 0..d {
for e_ax in 0..d {
let expect = t_in.slice(s![.., .., a_ax, c_ax, e_ax]).to_owned().dot(&q);
for i in 0..t_red.shape()[0] {
for j in 0..r {
assert_abs_diff_eq!(
t_red[[i, j, a_ax, c_ax, e_ax]],
expect[[i, j]],
epsilon = 1e-12
);
}
}
}
}
}
}
#[test]
pub(crate) fn production_builder_circle_reduces_rank_and_completes_stage1_step0_in_budget() {
let n_obs = 6usize;
let m = 5usize;
let d = 1usize;
let p = 2usize;
let k_atoms = 1usize;
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(m).unwrap());
let coords = array![[0.1], [0.45], [0.8], [0.1], [0.45], [0.8]];
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let mut basis_values = Array3::<f64>::zeros((k_atoms, n_obs, m));
basis_values.slice_mut(s![0, .., ..]).assign(&phi);
let mut basis_jacobian = Array4::<f64>::zeros((k_atoms, n_obs, m, d));
basis_jacobian.slice_mut(s![0, .., .., ..]).assign(&jet);
let mut decoder = Array3::<f64>::zeros((k_atoms, m, p));
decoder.slice_mut(s![0, .., ..]).assign(&array![
[0.05, -0.02],
[-0.05, 0.03],
[0.05, 0.01],
[0.02, -0.04],
[-0.02, 0.02]
]);
let mut penalties = Array3::<f64>::zeros((k_atoms, m, m));
penalties
.slice_mut(s![0, .., ..])
.assign(&Array2::<f64>::eye(m));
let logits = Array2::<f64>::zeros((n_obs, k_atoms));
let evaluators: Vec<Option<Arc<dyn SaeBasisSecondJet>>> = vec![Some(evaluator)];
let mut term = term_from_padded_blocks_with_mode(
n_obs,
p,
&[SaeAtomBasisKind::Periodic],
basis_values.view(),
basis_jacobian.view(),
&[m],
&[d],
decoder.view(),
penalties.view(),
logits.view(),
std::slice::from_ref(&coords),
AssignmentMode::ordered_beta_bernoulli(1.0, 1.0, false),
&evaluators,
)
.unwrap();
assert!(
term.atoms[0].basis_second_jet.is_some(),
"production builder must install the analytic second-jet evaluator so the \
#1117 rank-revealing reduction can fire",
);
let target = array![
[1.0, 0.0],
[0.0, 1.0],
[-1.0, 0.0],
[1.0, 0.0],
[0.0, 1.0],
[-1.0, 0.0]
];
let mut rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1)]);
let loss0 = term.loss(target.view(), &rho).unwrap().total();
let loss = term
.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 8, 0.05, 1.0e-3, 1.0e-3)
.unwrap();
assert_eq!(
term.atoms[0].basis_size(),
3,
"the rank-deficient circle must be reparametrized onto its r = 3 \
data-supported subspace at fit entry",
);
assert!(
loss.total().is_finite(),
"rank-deficient circle fit must return a finite loss, not stall: {}",
loss.total(),
);
assert!(
loss.total() <= loss0 + 1.0e-8,
"the joint fit must not increase the loss (loss0={loss0}, loss={})",
loss.total(),
);
assert!(
term.assignment.coords[0]
.as_flat()
.iter()
.all(|v| v.is_finite()),
"fitted coordinates must stay finite",
);
}
#[test]
pub(crate) fn rank_reduction_is_idempotent_on_already_reduced_atom() {
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(5).unwrap());
let coords = array![[0.1], [0.45], [0.8], [0.1], [0.45], [0.8]];
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let penalty = Array2::<f64>::eye(5);
let decoder = array![[0.05], [-0.05], [0.05], [0.02], [-0.02]];
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
penalty,
)
.unwrap()
.with_basis_second_jet(evaluator.clone());
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::ones((6, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
term.reduce_atoms_to_data_supported_rank().unwrap();
assert_eq!(term.atoms[0].basis_size(), 3);
let design_after_first = term.atoms[0].basis_values.clone();
let decoder_after_first = term.atoms[0].decoder_coefficients().clone();
term.reduce_atoms_to_data_supported_rank().unwrap();
assert_eq!(
term.atoms[0].basis_size(),
3,
"a second reduction pass on a full-rank reduced atom must be a no-op",
);
let design_after_second = &term.atoms[0].basis_values;
for i in 0..design_after_first.nrows() {
for j in 0..3 {
assert_eq!(
design_after_second[[i, j]],
design_after_first[[i, j]],
"idempotent reduction must leave the reduced design byte-identical",
);
}
}
let decoder_after_second = term.atoms[0].decoder_coefficients();
for i in 0..3 {
assert_eq!(
decoder_after_second[[i, 0]],
decoder_after_first[[i, 0]],
"idempotent reduction must leave the reduced decoder byte-identical",
);
}
}
#[test]
pub(crate) fn full_rank_circle_design_keeps_full_harmonic_depth_unchanged() {
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(5).unwrap());
let coords = array![[0.05], [0.27], [0.46], [0.68], [0.91]];
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let penalty = Array2::<f64>::eye(5);
let decoder = array![[0.05], [-0.05], [0.05], [0.02], [-0.02]];
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi.clone(),
jet,
decoder.clone(),
penalty,
)
.unwrap()
.with_basis_second_jet(evaluator.clone());
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::ones((5, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
term.reduce_atoms_to_data_supported_rank().unwrap();
assert_eq!(
term.atoms[0].basis_size(),
5,
"a full-rank circle design must keep all 5 harmonic columns",
);
let after_phi = &term.atoms[0].basis_values;
for i in 0..5 {
for j in 0..5 {
assert_eq!(
after_phi[[i, j]],
phi[[i, j]],
"full-rank basis must be unchanged by the (no-op) reduction",
);
}
}
let after = term.atoms[0].decoder_coefficients();
for i in 0..5 {
assert_eq!(
after[[i, 0]],
decoder[[i, 0]],
"full-rank decoder must be unchanged by the (no-op) reduction",
);
}
}
#[test]
pub(crate) fn solve_newton_step_escalates_ridge_on_non_pd_row_block() {
let coords = array![[0.1], [0.4], [0.7]];
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.05], [-0.05], [0.05]],
Array2::<f64>::zeros((3, 3)),
)
.unwrap();
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((3, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.20], [-0.10], [0.45]];
let rho = SaeManifoldRho::new(0.0, -20.0, vec![Array1::<f64>::zeros(1)]);
let result = term.solve_newton_step(target.view(), &rho, None, 1.0e-6, 1.0e-6);
assert!(
result.is_ok(),
"solve_newton_step should recover from degenerate H_tt via LM ridge escalation; got: {result:?}",
);
}
#[test]
pub(crate) fn sae_arrow_schur_beta_quadratic_model_matches_penalized_loss_change() {
let coords = array![[0.10], [0.35], [0.80]];
let (phi, jet) = periodic_basis(&coords);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.65], [-0.45], [0.25]],
array![[3.0, 0.4, -0.2], [0.1, 2.5, 0.3], [-0.5, 0.2, 1.8]],
)
.unwrap();
let assignment = SaeAssignment::from_blocks_with_mode(
Array2::<f64>::zeros((3, 1)),
vec![coords],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.20], [-0.10], [0.45]];
let rho = SaeManifoldRho::new(0.0, 1.3_f64.ln(), vec![array![0.9_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let beta0 = term.flatten_beta();
let loss0 = term.loss(target.view(), &rho).unwrap().total();
let mut direction = sys.gb.mapv(|v| -v);
let direction_norm = direction.iter().map(|v| v * v).sum::<f64>().sqrt();
assert!(direction_norm > 1.0e-12);
for value in direction.iter_mut() {
*value /= direction_norm;
}
let epsilon = 1.0e-3;
let delta = direction.mapv(|v| epsilon * v);
let beta_trial = beta0 + δ
term.set_flat_beta(beta_trial.view()).unwrap();
let actual = term.loss(target.view(), &rho).unwrap().total() - loss0;
let linear = sys.gb.dot(&delta);
let mut hbb_delta = Array1::<f64>::zeros(delta.len());
{
let op = sys.effective_penalty_op();
let d_slice = delta.as_slice().expect("delta is contiguous");
let hd_slice = hbb_delta.as_slice_mut().expect("hbb_delta is contiguous");
op.matvec(d_slice, hd_slice);
}
let quadratic = 0.5 * delta.dot(&hbb_delta);
let predicted = linear + quadratic;
let error = (actual - predicted).abs();
assert!(
error <= 1.0e-4,
"actual={actual:.12e}, predicted={predicted:.12e}, error={error:.12e}"
);
}
#[test]
pub(crate) fn separation_barrier_deferred_curvature_matches_dense_hbb_1610() {
let (term, _target, _rho) = small_two_atom_periodic_term();
let beta_dim = term.beta_dim();
let mut dense = ArrowSchurSystem::new(0, 0, beta_dim);
dense.gb = Array1::<f64>::zeros(beta_dim);
dense.hbb = Array2::<f64>::zeros((beta_dim, beta_dim));
let mut dense_atom_curv = vec![0.0_f64; term.k_atoms()];
let mut dense_curvature = Vec::new();
assert!(
term.add_sae_separation_barrier(
&mut dense,
1.0,
true,
&mut dense_atom_curv,
&mut dense_curvature
),
"fixture must activate the co-collapse separation barrier on the dense path"
);
assert!(
dense_atom_curv.iter().all(|v| *v == 0.0),
"dense path writes curvature directly to hbb, not the deferred atom accumulator"
);
assert!(
dense_curvature.is_empty(),
"dense path expands the curvature straight into hbb, not the deferred carrier"
);
let mut deferred = ArrowSchurSystem::new(0, 0, beta_dim);
deferred.gb = Array1::<f64>::zeros(beta_dim);
deferred.hbb = Array2::<f64>::zeros((0, 0));
let mut atom_curv = vec![0.0_f64; term.k_atoms()];
let mut sep_curvature = Vec::new();
assert!(
term.add_sae_separation_barrier(
&mut deferred,
1.0,
false,
&mut atom_curv,
&mut sep_curvature
),
"fixture must activate the co-collapse separation barrier on the deferred path"
);
for idx in 0..beta_dim {
assert!(
(dense.gb[idx] - deferred.gb[idx]).abs() <= 1.0e-12,
"dense and deferred paths must assemble the same barrier gradient at β[{idx}]"
);
}
let offsets = term.beta_offsets();
assert!(
!sep_curvature.is_empty(),
"deferred path must export the exact Gauss-Newton curvature carrier"
);
let mut deferred_hbb = Array2::<f64>::zeros((beta_dim, beta_dim));
for atom_idx in 0..term.k_atoms() {
let start = offsets[atom_idx];
let end = if atom_idx + 1 < offsets.len() {
offsets[atom_idx + 1]
} else {
beta_dim
};
assert!(
atom_curv[atom_idx] > 0.0,
"deferred path must export positive per-atom collapse-prevention curvature"
);
for idx in start..end {
deferred_hbb[[idx, idx]] += atom_curv[atom_idx];
}
}
for curvature in &sep_curvature {
deferred_hbb += &curvature.as_full_beta_op(beta_dim, &offsets).to_dense();
}
let scale = dense
.hbb
.iter()
.fold(0.0_f64, |acc, value| acc.max(value.abs()));
let tolerance = 1.0e-12 * (1.0 + scale);
for i in 0..beta_dim {
for j in 0..beta_dim {
assert!(
(dense.hbb[[i, j]] - deferred_hbb[[i, j]]).abs() <= tolerance,
"dense hbb and reconstructed deferred (ridge + factored curvature) must \
match at ({i},{j}): dense={} deferred={}",
dense.hbb[[i, j]],
deferred_hbb[[i, j]]
);
}
}
}
#[test]
pub(crate) fn sae_row_layout_from_topk_gates_is_exact() {
let coord_dims = vec![2usize, 1, 2];
let coord_offsets_full = vec![0usize, 2, 3];
let assignments = vec![
Array1::from_vec(vec![1.0, 0.0, 1.0]),
Array1::from_vec(vec![1.0, 1.0, 0.0]),
];
let layout =
SaeRowLayout::from_topk_gates(&assignments, 2, coord_dims, coord_offsets_full).unwrap();
assert_eq!(layout.active_atoms[0], vec![0, 2]);
assert_eq!(layout.active_atoms[1], vec![0, 1]);
assert_eq!(layout.row_q_active(0), 4);
assert_eq!(layout.row_q_active(1), 3);
let compact = vec![1.0_f64, 2.0, 3.0, 4.0];
let mut full = vec![0.0_f64; 5];
layout.expand_row(0, &compact, &mut full);
assert_eq!(full[0], 1.0);
assert_eq!(full[1], 2.0);
assert_eq!(full[2], 0.0);
assert_eq!(full[3], 3.0);
assert_eq!(full[4], 4.0);
}
#[test]
pub(crate) fn from_topk_gates_large_k_support_is_exact() {
let (k_atoms, d, k_true, n) = (100_000_usize, 1, 4, 4);
let planted: Vec<usize> = (0..k_true).map(|j| j * k_atoms / k_true).collect();
let assignments: Vec<Array1<f64>> = (0..n)
.map(|_| {
let mut a = vec![0.0_f64; k_atoms];
for &atom in &planted {
a[atom] = 1.0;
}
Array1::from_vec(a)
})
.collect();
let coord_offsets: Vec<usize> = (0..k_atoms).collect();
let layout =
SaeRowLayout::from_topk_gates(&assignments, k_true, vec![d; k_atoms], coord_offsets)
.unwrap();
for row in 0..n {
assert_eq!(layout.active_atoms[row], planted, "row {row} wrong atoms");
assert_eq!(layout.row_q_active(row), k_true * d);
}
let compact_work: usize = (0..n).map(|r| layout.row_q_active(r).pow(2)).sum();
assert!(compact_work < n * (k_atoms * d).pow(2) / 1_000_000);
}
#[test]
pub(crate) fn sae_row_layout_from_topk_gates_large_k_work_scales_with_support() {
let n = 4usize;
let k_atoms = 100_000usize;
let cap = 8usize;
let mut planted: Vec<Vec<usize>> = Vec::with_capacity(n);
let mut assignments: Vec<Array1<f64>> = Vec::with_capacity(n);
for row in 0..n {
let mut a = Array1::<f64>::zeros(k_atoms);
let mut plant = Vec::with_capacity(cap);
for j in 0..cap {
let idx = (row + j * (k_atoms / cap)) % k_atoms;
a[idx] = 1.0;
plant.push(idx);
}
plant.sort_unstable();
planted.push(plant);
assignments.push(a);
}
let coord_dims = vec![1usize; k_atoms];
let coord_offsets_full: Vec<usize> = (0..k_atoms).collect();
let layout =
SaeRowLayout::from_topk_gates(&assignments, cap, coord_dims, coord_offsets_full).unwrap();
for row in 0..n {
assert_eq!(
layout.active_atoms[row], planted[row],
"row {row}: support recovery mismatch"
);
assert_eq!(layout.row_q_active(row), cap, "row {row}: q_active");
}
let compact_work: usize = (0..n)
.map(|row| {
let q = layout.row_q_active(row);
q * q
})
.sum();
assert_eq!(compact_work, n * cap * cap);
let dense_q = k_atoms;
let dense_work = n * dense_q * dense_q;
let work_ratio = (k_atoms / cap) * (k_atoms / cap);
assert_eq!(
dense_work / compact_work,
work_ratio,
"compact row-layout work must be exactly (K/cap)² below dense full-K work"
);
assert!(
work_ratio >= 100_000_000,
"the compact path must be astronomically (≥1e8×) cheaper than dense full-K"
);
}
#[test]
pub(crate) fn fixed_decoder_assembly_skips_beta_tier_1407() {
let (mut term, target, rho) = small_two_atom_periodic_term();
let full = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("full joint assemble_arrow_schur");
assert!(
full.penalty_op.is_some(),
"full joint assembly must install the β-tier curvature operator \
(matrix-free smoothness + G⊗I data-fit block)"
);
assert_eq!(
full.k,
term.beta_dim(),
"full joint assembly must carry the full β-tier width"
);
assert!(
full.gb.len() == term.beta_dim() && full.gb.iter().all(|v| v.is_finite()),
"full joint assembly must carry a finite β-tier gradient gb of width beta_dim"
);
let n_rows = full.rows.len();
term.fixed_decoder_assembly = true;
let fixed = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("fixed-decoder assemble_arrow_schur");
term.fixed_decoder_assembly = false;
assert!(
fixed.penalty_op.is_none(),
"fixed-decoder assembly must install NO β-tier curvature operator (the β \
tier is dead work when the decoder is frozen)"
);
assert_eq!(
fixed.hbb.dim(),
(0, 0),
"fixed-decoder assembly must build NO dense β-Hessian (the β tier is \
dead work when the decoder is frozen); got hbb {:?}",
fixed.hbb.dim()
);
assert_eq!(
fixed.rows.len(),
n_rows,
"fixed-decoder assembly must keep every per-row htt/gt block"
);
for (i, row) in fixed.rows.iter().enumerate() {
assert!(
row.htt.iter().all(|v| v.is_finite()) && row.gt.iter().all(|v| v.is_finite()),
"fixed-decoder row {i} htt/gt must be finite"
);
assert_eq!(
row.htt.dim().0,
row.gt.len(),
"fixed-decoder row {i} htt must be square and match gt length"
);
assert!(
row.gt.len() > 0,
"fixed-decoder row {i} must carry a non-empty latent block"
);
}
}
#[test]
pub(crate) fn sae_mechsparsity_beta_block_routes_through_arrow_schur_gb() {
let coords = array![[0.10], [0.35], [0.80]];
let (phi, jet) = periodic_basis(&coords);
let decoder = array![
[0.7, -0.2, 0.05, 0.4],
[-0.5, 0.6, -0.1, 0.3],
[0.2, 0.0, -0.4, -0.6],
];
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder.clone(),
Array2::<f64>::eye(3),
)
.unwrap();
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((3, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let m = 3usize;
let p = 4usize;
let slice = PsiSlice::full(m * p, Some(m));
let penalty = MechanismSparsityPenalty::new(
slice,
vec![vec![0, 1], vec![2, 3]],
1.0,
1.0e-6,
(term.n_obs()) as f64,
false,
)
.unwrap();
let mut registry = AnalyticPenaltyRegistry::new();
registry.push(AnalyticPenaltyKind::MechanismSparsity(Arc::new(penalty)));
let target = array![
[0.20, 0.10, -0.05, 0.25],
[-0.10, 0.30, 0.15, -0.20],
[0.45, -0.05, 0.10, 0.30],
];
let rho = SaeManifoldRho::new(0.0, -6.0, vec![Array1::<f64>::zeros(1)]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, Some(®istry))
.unwrap();
assert_eq!(sys.gb.len(), m * p, "gb should match flatten_beta length");
let mut absmax = 0.0_f64;
for v in sys.gb.iter().copied() {
assert!(v.is_finite());
if v.abs() > absmax {
absmax = v.abs();
}
}
assert!(
absmax > 1.0e-6,
"MechSparsity must inject a non-trivial gradient into the SAE arrow-Schur gb; absmax={absmax:.3e}"
);
let sys_no_penalty = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let beta = term.flatten_beta();
let expected = {
let s = (0.5_f64.powi(2) + 0.6_f64.powi(2) + 1.0e-12).sqrt();
(2.0_f64).sqrt() * (-0.5_f64) / s
};
let delta = sys.gb[1 * p + 0] - sys_no_penalty.gb[1 * p + 0];
assert!(
(delta - expected).abs() <= 1.0e-6,
"expected MechSparsity gb contribution at (basis=1, feat=0) ≈ {expected:.6e}, \
got Δgb={delta:.6e} (gb_with={:.6e}, gb_without={:.6e}, beta entry = {})",
sys.gb[1 * p + 0],
sys_no_penalty.gb[1 * p + 0],
beta[1 * p + 0]
);
}
pub(crate) fn smoothed_nuclear_norm(decoder: &Array2<f64>, eps: f64) -> f64 {
let (_u, s, _vt) = decoder
.clone()
.svd(false, false)
.expect("the decoder is finite, so its SVD converges");
s.iter()
.map(|sigma| (sigma * sigma + eps * eps).sqrt() - eps)
.sum()
}
#[test]
pub(crate) fn sae_nuclear_norm_beta_block_routes_through_gb_and_shrinks_spectrum() {
let coords = array![[0.10], [0.35], [0.80]];
let (phi, jet) = periodic_basis(&coords);
let decoder = array![
[0.9, -0.2, 0.05, 0.4],
[-0.5, 0.7, -0.1, 0.3],
[0.2, 0.1, -0.8, -0.6],
];
let m = 3usize;
let p = 4usize;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder.clone(),
Array2::<f64>::eye(3),
)
.unwrap();
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((3, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let eps = 1.0e-6;
let slice = PsiSlice::full(m * p, Some(m));
let penalty = NuclearNormPenalty::new(slice, 1.0, p, eps, None, false).unwrap();
let mut registry = AnalyticPenaltyRegistry::new();
registry.push(AnalyticPenaltyKind::NuclearNorm(Arc::new(penalty)));
term.validate_analytic_penalty_registry(®istry)
.expect("NuclearNorm must be accepted (redirected to the β block)");
let target = array![
[0.20, 0.10, -0.05, 0.25],
[-0.10, 0.30, 0.15, -0.20],
[0.45, -0.05, 0.10, 0.30],
];
let rho = SaeManifoldRho::new(0.0, -6.0, vec![Array1::<f64>::zeros(1)]);
let baseline = term
.assemble_arrow_schur(target.view(), &rho, None)
.unwrap();
let sys = term
.assemble_arrow_schur(target.view(), &rho, Some(®istry))
.unwrap();
assert_eq!(sys.gb.len(), m * p, "gb should match flatten_beta length");
assert_eq!(
baseline.gb.len(),
m * p,
"baseline gb should match flatten_beta length"
);
let mut absmax = 0.0_f64;
let mut penalty_grad = Array1::<f64>::zeros(m * p);
for ((dst, sys_g), baseline_g) in penalty_grad
.iter_mut()
.zip(sys.gb.iter())
.zip(baseline.gb.iter())
{
let v = *sys_g - *baseline_g;
assert!(v.is_finite());
*dst = v;
absmax = absmax.max(v.abs());
}
assert!(
absmax > 1.0e-6,
"NuclearNorm must inject a non-trivial gradient into the SAE \
arrow-Schur gb; absmax={absmax:.3e}"
);
let per_atom = NuclearNormPenalty::new(
PsiSlice {
range: 0..m * p,
latent_dim: Some(p),
},
1.0,
m,
eps,
None,
false,
)
.unwrap();
let beta = term.flatten_beta();
let ref_grad = per_atom.grad_target(beta.view(), Array1::<f64>::zeros(0).view());
for j in 0..m * p {
assert!(
(penalty_grad[j] - ref_grad[j]).abs() <= 1.0e-9,
"penalty gb[{j}]={:.12e} must equal analytic spectral grad {:.12e}",
penalty_grad[j],
ref_grad[j]
);
}
let base_norm = smoothed_nuclear_norm(&decoder, eps);
let step = 1.0e-2;
let mut shrunk = decoder.clone();
for ((row, feat), value) in shrunk.indexed_iter_mut() {
*value -= step * penalty_grad[row * p + feat];
}
let shrunk_norm = smoothed_nuclear_norm(&shrunk, eps);
assert!(
shrunk_norm < base_norm,
"a step along gb must shrink the decoder spectrum: \
before={base_norm:.9e}, after={shrunk_norm:.9e}"
);
assert!(sys.hbb.is_empty());
let mut hbb_diag = vec![0.0_f64; m * p];
sys.effective_penalty_op().diagonal(&mut hbb_diag);
for i in 0..m * p {
assert!(
hbb_diag[i] >= -1.0e-9,
"hbb diagonal must be non-negative (PSD majorizer); hbb[{i},{i}]={:.3e}",
hbb_diag[i]
);
}
}
fn hetero_compat_term(d0: usize, d1: usize) -> SaeManifoldTerm {
let n = 4usize;
let p = 3usize;
let m = 2usize;
let make_atom = |name: &'static str, d: usize| {
SaeManifoldAtom::new_with_provided_function_gram(
name,
SaeAtomBasisKind::EuclideanPatch,
d,
Array2::<f64>::ones((n, m)),
Array3::<f64>::zeros((n, m, d)),
Array2::<f64>::zeros((m, p)),
Array2::<f64>::eye(m),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction")
};
let manifold = |d: usize| {
if d == 1 {
LatentManifold::Euclidean
} else {
LatentManifold::Product(vec![LatentManifold::Euclidean; d])
}
};
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 2)),
vec![Array2::<f64>::zeros((n, d0)), Array2::<f64>::zeros((n, d1))],
vec![manifold(d0), manifold(d1)],
AssignmentMode::softmax(1.0),
)
.expect(
"assignment fixture: logits, coordinate blocks and manifolds agree in block count and rows",
);
SaeManifoldTerm::new(
vec![make_atom("atom0", d0), make_atom("atom1", d1)],
assignment,
)
.expect("term fixture: every atom's row count matches the assignment's")
}
#[test]
pub(crate) fn validate_heterogeneous_atom_compatibility_covers_registry_and_native_ard() {
use gam_terms::analytic_penalties::{
BlockOrthogonalityPenalty, IsometryPenalty, PenaltyConcavity, ScadMcpPenalty,
};
let hetero = hetero_compat_term(2, 1);
let empty_registry = AnalyticPenaltyRegistry::new();
let mut structural_registry = AnalyticPenaltyRegistry::new();
structural_registry.push(AnalyticPenaltyKind::BlockOrthogonality(Arc::new(
BlockOrthogonalityPenalty::new(
PsiSlice::full(4 * 2, Some(2)),
vec![vec![0], vec![1]],
1.0,
4,
false,
)
.unwrap(),
)));
let err = hetero
.validate_heterogeneous_atom_compatibility(Some(&structural_registry), false)
.expect_err("heterogeneous dims + a fixed-d structural penalty must be refused");
assert!(
err.contains("heterogeneous atom coordinate dims"),
"message must name the heterogeneous conflict: {err}"
);
assert!(
err.contains("uniform atom_dim"),
"message must name the uniform-dims resolution: {err}"
);
assert!(
err.contains("block_orthogonality"),
"message must name the offending penalty kind: {err}"
);
let mut iso_registry = AnalyticPenaltyRegistry::new();
iso_registry.push(AnalyticPenaltyKind::Isometry(Arc::new(
IsometryPenalty::new_euclidean(PsiSlice::full(4 * 2, Some(2)), 2),
)));
hetero
.validate_heterogeneous_atom_compatibility(Some(&iso_registry), false)
.expect("isometry gauge composes per atom on a heterogeneous dictionary");
let mut scad_registry = AnalyticPenaltyRegistry::new();
scad_registry.push(AnalyticPenaltyKind::ScadMcp(Arc::new(
ScadMcpPenalty::new(
PsiSlice::full(4 * 2, Some(2)),
1.0,
4,
3.7,
1.0e-3,
PenaltyConcavity::Scad,
false,
)
.unwrap(),
)));
hetero
.validate_heterogeneous_atom_compatibility(Some(&scad_registry), false)
.expect("element-wise SCAD-MCP composes on a heterogeneous dictionary");
hetero
.validate_heterogeneous_atom_compatibility(Some(&empty_registry), true)
.expect("native ARD composes per atom on a heterogeneous dictionary");
hetero
.validate_heterogeneous_atom_compatibility(None, true)
.expect("native ARD (no registry) composes per atom on a heterogeneous dictionary");
hetero
.validate_heterogeneous_atom_compatibility(Some(&iso_registry), true)
.expect("ARD + isometry gauge compose together on a heterogeneous dictionary");
let homo = hetero_compat_term(2, 2);
homo.validate_heterogeneous_atom_compatibility(Some(&structural_registry), true)
.expect("homogeneous coord dims dispatch every row-block penalty cleanly");
hetero
.validate_heterogeneous_atom_compatibility(Some(&empty_registry), false)
.expect("heterogeneous dims with no row-block penalty and no ARD is admitted");
hetero
.validate_heterogeneous_atom_compatibility(None, false)
.expect("heterogeneous dims with no registry and no ARD is admitted");
}
fn ard_atom_and_coord(
name: &'static str,
manifold: LatentManifold,
coords: Array2<f64>,
) -> (SaeManifoldAtom, LatentManifold, Array2<f64>) {
let (n, d) = coords.dim();
let m = 2usize;
let p = 3usize;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
name,
SaeAtomBasisKind::EuclideanPatch,
d,
Array2::<f64>::ones((n, m)),
Array3::<f64>::zeros((n, m, d)),
Array2::<f64>::zeros((m, p)),
Array2::<f64>::eye(m),
)
.expect("atom fixture: basis, jet, decoder and Gram shapes agree by construction");
(atom, manifold, coords)
}
#[test]
pub(crate) fn periodic_ard_centered_bessel_value_gradient_survive_domain_edge() {
let n = 3usize;
let coords = Array2::<f64>::zeros((n, 1));
let manifold = LatentManifold::Circle { period: 1.0 };
let (atom, manifold, coords) = ard_atom_and_coord("circle", manifold, coords);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![manifold],
AssignmentMode::softmax(1.0),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let log_alpha = LOG_STRENGTH_MAX - 1.0e-3;
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![log_alpha]]);
let value = term.ard_value(&rho).unwrap();
let log_eta = log_alpha - 2.0 * std::f64::consts::TAU.ln();
let expected = -0.5 * n as f64 * (std::f64::consts::TAU.ln() + log_eta);
assert!(value.is_finite());
assert!(
(value - expected).abs() < 2.0e-8,
"centered periodic normalizer drifted: value={value}, expected={expected}"
);
let derivative = term.ard_log_precision_explicit_derivatives(&rho).unwrap()[0][0];
assert!(derivative.is_finite());
assert!(
(derivative + 0.5 * n as f64).abs() < 1.0e-12,
"periodic normalizer derivative lost its -n/2 limit: {derivative}"
);
let step = 1.0e-4;
let plus = SaeManifoldRho::new(0.0, 0.0, vec![array![log_alpha + step]]);
let minus = SaeManifoldRho::new(0.0, 0.0, vec![array![log_alpha - step]]);
let finite_difference =
(term.ard_value(&plus).unwrap() - term.ard_value(&minus).unwrap()) / (2.0 * step);
assert!(
(finite_difference - derivative).abs() < 2.0e-7,
"periodic ARD value/gradient mismatch: analytic={derivative}, finite_difference={finite_difference}"
);
}
#[test]
pub(crate) fn native_ard_energy_composes_additively_on_mixed_dictionary() {
let n = 3usize;
let circle_coords = array![[0.1_f64], [0.6], [0.9]];
let patch_coords = array![[0.2_f64, -0.3], [0.5, 0.4], [-0.1, 0.7]];
let linear_coords = array![[1.0_f64], [-0.5], [0.3]];
let circle_m = LatentManifold::Circle { period: 1.0 };
let patch_m = LatentManifold::Product(vec![LatentManifold::Euclidean; 2]);
let linear_m = LatentManifold::Euclidean;
let ard0 = array![0.2_f64];
let ard1 = array![0.1_f64, -0.3];
let ard2 = array![-0.15_f64];
let (a0, m0, c0) = ard_atom_and_coord("circle", circle_m.clone(), circle_coords.clone());
let (a1, m1, c1) = ard_atom_and_coord("patch", patch_m.clone(), patch_coords.clone());
let (a2, m2, c2) = ard_atom_and_coord("linear", linear_m.clone(), linear_coords.clone());
let joint_assign = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 3)),
vec![c0, c1, c2],
vec![m0, m1, m2],
AssignmentMode::softmax(1.0),
)
.unwrap();
let joint = SaeManifoldTerm::new(vec![a0, a1, a2], joint_assign).unwrap();
let joint_rho = SaeManifoldRho::new(0.0, 0.0, vec![ard0.clone(), ard1.clone(), ard2.clone()]);
let joint_energy = joint.ard_value(&joint_rho).unwrap();
let single_energy = |name: &'static str,
manifold: LatentManifold,
coords: Array2<f64>,
ard: Array1<f64>|
-> f64 {
let (atom, m, c) = ard_atom_and_coord(name, manifold, coords);
let assign = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![c],
vec![m],
AssignmentMode::softmax(1.0),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom], assign).unwrap();
let rho = SaeManifoldRho::new(0.0, 0.0, vec![ard]);
term.ard_value(&rho).unwrap()
};
let sum = single_energy("circle", circle_m, circle_coords, ard0)
+ single_energy("patch", patch_m, patch_coords, ard1)
+ single_energy("linear", linear_m, linear_coords, ard2);
assert!(
(joint_energy - sum).abs() < 1.0e-12,
"native ARD must compose additively over a mixed dictionary: \
joint={joint_energy}, per-atom sum={sum}"
);
assert!(
joint_energy.abs() > 1.0e-6,
"mixed-dictionary ARD energy should be non-trivial, got {joint_energy}"
);
}
#[derive(Debug)]
pub(crate) struct TestPeriodicEvaluator;
impl SaeBasisEvaluator for TestPeriodicEvaluator {
fn second_jet_dyn(&self, coords: ArrayView2<'_, f64>) -> Option<Result<Array4<f64>, String>> {
if coords.ncols() != 1 {
return Some(Err(format!(
"TestPeriodicEvaluator::second_jet_dyn: expected latent_dim 1, got {}",
coords.ncols()
)));
}
let n = coords.nrows();
let two_pi = 2.0 * std::f64::consts::PI;
let freq2 = two_pi * two_pi;
let mut h = Array4::<f64>::zeros((n, 3, 1, 1));
for row in 0..n {
let angle = two_pi * coords[[row, 0]];
h[[row, 1, 0, 0]] = -freq2 * angle.sin();
h[[row, 2, 0, 0]] = -freq2 * angle.cos();
}
Some(Ok(h))
}
fn third_jet_dyn(&self, coords: ArrayView2<'_, f64>) -> Option<Result<Array5<f64>, String>> {
if coords.ncols() != 1 {
return Some(Err(format!(
"TestPeriodicEvaluator::third_jet_dyn: expected latent_dim 1, got {}",
coords.ncols()
)));
}
let n = coords.nrows();
let two_pi = 2.0 * std::f64::consts::PI;
let freq3 = two_pi * two_pi * two_pi;
let mut h = Array5::<f64>::zeros((n, 3, 1, 1, 1));
for row in 0..n {
let angle = two_pi * coords[[row, 0]];
h[[row, 1, 0, 0, 0]] = -freq3 * angle.cos();
h[[row, 2, 0, 0, 0]] = freq3 * angle.sin();
}
Some(Ok(h))
}
fn evaluate(&self, coords: ArrayView2<'_, f64>) -> Result<(Array2<f64>, Array3<f64>), String> {
Ok(periodic_basis(&coords.to_owned()))
}
}
#[test]
pub(crate) fn criterion_lane_gap_is_exactly_the_evidence_logdet_gap_2509() {
use crate::manifold::construction::StreamingRankInputs;
let (term0, target, rho) = small_two_atom_periodic_term_at_shared_inner_state();
let mut dense = term0.clone();
let mut streaming = term0;
let (dense_cost, dense_loss, cache) = dense
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
FROZEN_INNER_STATE,
0.25,
1.0e-4,
1.0e-4,
)
.expect("dense criterion must evaluate at the pinned #2509 witness state");
let (log_a, log_a_tt) = dense
.exact_observed_information_log_dets(&rho, target.view(), &cache)
.expect("exact observed-information log-dets");
let (stream_cost, stream_loss) = streaming
.penalized_quasi_laplace_criterion_streaming_exact(
target.view(),
&rho,
None,
FROZEN_INNER_STATE,
0.25,
1.0e-4,
1.0e-4,
)
.expect("streaming criterion must evaluate at the pinned #2509 witness state");
let mut rank_inputs = StreamingRankInputs::default();
let log_b = streaming
.streaming_exact_arrow_log_det(target.view(), &rho, None, Some(&mut rank_inputs))
.expect("streaming majorizer log-det");
let log_b_tt = rank_inputs.log_det_tt;
let measured_gap = dense_cost - stream_cost;
let logdet_gap = 0.5 * ((log_a - log_a_tt) - (log_b - log_b_tt));
let scale = [
dense_cost.abs(),
stream_cost.abs(),
log_a.abs(),
log_a_tt.abs(),
log_b.abs(),
log_b_tt.abs(),
dense_loss.total().abs(),
]
.into_iter()
.fold(1.0_f64, f64::max);
let tolerance = 4096.0 * f64::EPSILON * scale;
assert!(
(dense_loss.total() - stream_loss.total()).abs() <= tolerance,
"#2509: the two lanes' converged losses must agree before any log-det \
claim is meaningful: dense {} vs streaming {} (tolerance {tolerance:.3e})",
dense_loss.total(),
stream_loss.total()
);
assert!(
(measured_gap - logdet_gap).abs() <= tolerance,
"#2509: the dense↔streaming criterion gap must be EXACTLY the evidence \
log-determinant gap, so a second desync cannot hide inside it.\n \
dense_cost = {dense_cost}\n streaming_cost = {stream_cost}\n \
measured_gap = {measured_gap}\n logdet_gap = {logdet_gap}\n \
residual = {:.6e} (tolerance {tolerance:.3e})\n \
log|A| = {log_a} log|A_tt| = {log_a_tt}\n \
log|B| = {log_b} log|B_tt| = {log_b_tt}",
measured_gap - logdet_gap
);
}
#[test]
pub(crate) fn inner_kkt_gate_is_extensive_while_the_intensive_certificate_exists_2681() {
let (mut term, target, rho) = small_two_atom_periodic_term();
let refusal = term.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
1,
0.25,
1.0e-4,
1.0e-4,
);
let error = match refusal {
Ok(_) => {
eprintln!(
"#2681: fixture now CONVERGES — the witnesses are unblocked, re-read the issue"
);
return;
}
Err(error) => error,
};
let message = error.numerical_message().unwrap_or_else(|| {
panic!(
"#2681: expected the numerical (inner non-convergence) refusal, got a \
different typed variant: {error:?}"
)
});
assert!(
message.contains("met tolerance"),
"#2681: expected the inner non-convergence refusal, got: {message}"
);
let system = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("the best-seen state must still assemble");
let raw_sq = SaeManifoldTerm::system_grad_norm_sq(&system);
let raw = raw_sq.sqrt();
let extensive_bound = SAE_MANIFOLD_INNER_GRAD_REL_TOL * term.inner_iterate_scale();
let scaled_max = SaeManifoldTerm::system_scaled_grad_max(&system)
.expect("the parameter-space certificate must resolve on this fixture");
let intensive_bound = SAE_MANIFOLD_INNER_GRAD_REL_TOL
* term
.inner_iterate_max()
.expect("the intensive iterate scale must resolve on this fixture");
eprintln!(
"#2681 certificates at the refused state:\n \
EXTENSIVE (the gate): ‖g‖₂ = {raw:.6e} bound = {extensive_bound:.6e} \
over = {:.3e}x\n \
INTENSIVE (parameter space, unwired here): scaled_grad_max = {scaled_max:.6e} \
bound = {intensive_bound:.6e} over = {:.3e}x",
raw / extensive_bound,
scaled_max / intensive_bound,
);
let extensive_over = raw / extensive_bound;
let intensive_over = scaled_max / intensive_bound;
assert!(
extensive_over > 10.0 && intensive_over > 10.0,
"#2681: a certificate has come within range at the refused state \
(extensive {extensive_over:.3e}x, intensive {intensive_over:.3e}x). The \
measurement that routed this to #2653 — both certificates refusing by \
orders of magnitude, so the bar was not the defect — no longer holds and \
the three #2509 witnesses should be re-checked."
);
assert!(
scaled_max > 1.0e-2,
"#2681: the parameter-space displacement has dropped to {scaled_max:.6e}. \
The finding that this state is nowhere near stationary IN PARAMETER UNITS \
(measured 4.23e-1) is what ruled out re-denominating the bar; if the solve \
now lands this close, that conclusion needs re-taking."
);
}