use super::tests::{
TestPeriodicEvaluator, diagonal_latent_cache, periodic_basis, warmstart_test_objective,
warmstart_test_objective_with_evaluator,
};
use super::*;
use crate::assignment::{AssignmentMode, SaeAssignment};
use approx::assert_abs_diff_eq;
use gam_terms::latent::LatentManifold;
use ndarray::array;
use std::sync::Arc;
pub(crate) fn euclidean_line_contract_fixture() -> (SaeManifoldTerm, Array2<f64>, SaeManifoldRho) {
let n = 150usize;
let p = 8usize;
let mut coords = Array2::<f64>::zeros((n, 1));
let mut z = Array2::<f64>::zeros((n, p));
for row in 0..n {
let u = -1.0 + 2.0 * row as f64 / (n as f64 - 1.0);
coords[[row, 0]] = 2.5 + 3.0 * u;
for col in 0..p {
let linear_loading = 0.35 + 0.07 * col as f64;
let offset = 0.08 * ((col % 3) as f64 - 1.0);
let phase = (row * (col + 3)) as f64;
let noise = 0.04 * (phase.sin() + 0.5 * (0.37 * phase).cos());
z[[row, col]] = offset + linear_loading * u + noise;
}
}
let evaluator = Arc::new(EuclideanPatchEvaluator::new(1, 2).expect("evaluator"));
let (phi, jet) = evaluator.evaluate(coords.view()).expect("basis");
let m = phi.ncols();
let smooth_penalty =
gam_terms::basis::create_difference_penalty_matrix(m, 2, None).expect("penalty");
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"contract-line",
SaeAtomBasisKind::EuclideanPatch,
1,
phi,
jet,
Array2::<f64>::zeros((m, p)),
smooth_penalty,
)
.expect("atom")
.with_basis_second_jet(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Euclidean],
AssignmentMode::softmax(1.0),
)
.expect("assignment");
let term = SaeManifoldTerm::new(vec![atom], assignment).expect("term");
let rho = SaeManifoldRho::new(0.0, (0.01_f64).ln(), vec![Array1::<f64>::zeros(1)]);
(term, z, rho)
}
pub(crate) fn assert_contract_close_with_floor(
label: &str,
analytic: f64,
finite_difference: f64,
fd_roundoff_floor: f64,
) {
let abs_diff = (analytic - finite_difference).abs();
let scale = finite_difference.abs().max(analytic.abs());
let tol = 1.0e-5 * scale + fd_roundoff_floor;
let rel = abs_diff / scale.max(1.0e-12);
assert!(
abs_diff <= tol,
"{label}: analytic={analytic:.12e} fd={finite_difference:.12e} \
rel={rel:.3e} abs_diff={abs_diff:.3e} tol={tol:.3e} \
(fd_roundoff_floor={fd_roundoff_floor:.3e})"
);
}
pub(crate) fn fd_roundoff_floor(f_plus: f64, f_minus: f64, h: f64) -> f64 {
const SAFETY: f64 = 16.0;
SAFETY * f64::EPSILON * f_plus.abs().max(f_minus.abs()) / (2.0 * h)
}
fn assert_decoder_gradient_matches_fd(
term: &mut SaeManifoldTerm,
z: &Array2<f64>,
rho: &SaeManifoldRho,
basis_col: usize,
out_col: usize,
p: usize,
h: f64,
) {
let sys = term
.assemble_arrow_schur(z.view(), rho, None)
.expect("decoder assemble");
assert_eq!(sys.k, term.beta_dim());
let beta_idx = basis_col * p + out_col;
let analytic = sys.gb[beta_idx];
let base = term.atoms[0].decoder_coefficients()[[basis_col, out_col]];
term.atoms[0].decoder_coefficients_mut()[[basis_col, out_col]] = base + h;
let f_plus = term
.penalized_objective_total(z.view(), rho, None, 1.0)
.expect("decoder f+");
term.atoms[0].decoder_coefficients_mut()[[basis_col, out_col]] = base - h;
let f_minus = term
.penalized_objective_total(z.view(), rho, None, 1.0)
.expect("decoder f-");
term.atoms[0].decoder_coefficients_mut()[[basis_col, out_col]] = base;
let fd = (f_plus - f_minus) / (2.0 * h);
assert_contract_close_with_floor(
&format!("CONTRACT decoder ({basis_col},{out_col})"),
analytic,
fd,
fd_roundoff_floor(f_plus, f_minus, h),
);
}
#[test]
pub(crate) fn euclidean_line_decoder_gradient_matches_penalized_objective_fd() {
let (mut term, z, mut rho) = euclidean_line_contract_fixture();
let p = term.output_dim();
let h = 1.0e-6;
assert_eq!(
term.atoms[0].basis_size(),
3,
"the degree-2 euclidean fixture must seed a width-3 [1,t,t²] basis"
);
for (basis_col, out_col) in [(0usize, 0usize), (1, 3), (2, 7)] {
assert_decoder_gradient_matches_fd(&mut term, &z, &rho, basis_col, out_col, p, h);
}
let ridge = 1.0e-6;
for step in 0..6 {
let loss = term
.run_joint_fit_arrow_schur(z.view(), &mut rho, None, 1, 1.0, ridge, ridge)
.unwrap_or_else(|err| panic!("warm step {step} failed: {err}"));
assert!(
loss.total().is_finite(),
"warm step {step} loss is non-finite"
);
}
let sys_coord = term
.assemble_arrow_schur(z.view(), &rho, None)
.expect("coord assemble");
assert_eq!(
sys_coord.k,
term.beta_dim(),
"p=8 contract fixture must stay on full-B coordinates"
);
assert!(
!term.frames_active(),
"p=8 contract fixture must not activate a frame"
);
let h = 1.0e-6;
for row in [3usize, 75, 140] {
let analytic = sys_coord.rows[row].gt[0];
let base_coord = term.assignment.coords[0].as_matrix()[[row, 0]];
let mut plus_coords = term.assignment.coords[0].as_matrix();
plus_coords[[row, 0]] = base_coord + h;
let plus_flat = Array1::from_iter(plus_coords.iter().copied());
term.assignment.coords[0].set_flat(plus_flat.view());
term.refresh_basis_from_current_coords()
.expect("plus refresh");
let f_plus = term
.penalized_objective_total(z.view(), &rho, None, 1.0)
.expect("coord f+");
let mut minus_coords = term.assignment.coords[0].as_matrix();
minus_coords[[row, 0]] = base_coord - h;
let minus_flat = Array1::from_iter(minus_coords.iter().copied());
term.assignment.coords[0].set_flat(minus_flat.view());
term.refresh_basis_from_current_coords()
.expect("minus refresh");
let f_minus = term
.penalized_objective_total(z.view(), &rho, None, 1.0)
.expect("coord f-");
let mut restored_coords = term.assignment.coords[0].as_matrix();
restored_coords[[row, 0]] = base_coord;
let restored_flat = Array1::from_iter(restored_coords.iter().copied());
term.assignment.coords[0].set_flat(restored_flat.view());
term.refresh_basis_from_current_coords()
.expect("restore refresh");
let fd = (f_plus - f_minus) / (2.0 * h);
assert_contract_close_with_floor(
&format!("CONTRACT coord row {row}"),
analytic,
fd,
fd_roundoff_floor(f_plus, f_minus, h),
);
}
let fitted_m = term.atoms[0].basis_size();
assert!(
fitted_m >= 1 && fitted_m <= 3,
"fitted euclidean basis width must be in [1,3]; got {fitted_m}"
);
for basis_col in 0..fitted_m {
for &out_col in &[0usize, p / 2, p - 1] {
assert_decoder_gradient_matches_fd(&mut term, &z, &rho, basis_col, out_col, p, h);
}
}
}
#[test]
fn sae_isometry_assembled_curvature_is_decoder_scale_invariant() {
use gam_terms::analytic_penalties::{
AnalyticPenaltyKind, AnalyticPenaltyRegistry, IsometryPenalty, PsiSlice,
};
let n = 24usize;
let p = 4usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.5) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let m = phi.ncols();
let base_decoder = Array2::from_shape_fn((m, p), |(b, c)| {
let scale = 1.0 / (1.0 + b as f64);
scale * ((b as f64 + 1.0) * (c as f64 + 1.0)).cos()
});
let isometry_curvature_norm = |lambda: f64| -> f64 {
let decoder = &base_decoder * lambda;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"iso_scale",
SaeAtomBasisKind::Periodic,
1,
phi.clone(),
jet.clone(),
decoder.clone(),
Array2::<f64>::eye(m),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let target = phi.dot(&decoder);
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),
)
.expect("fixture assignment: one logit column and one coord block per atom");
let mut term = SaeManifoldTerm::new(vec![atom], assignment)
.expect("fixture term: every atom's basis width matches its assignment block");
let mut registry = AnalyticPenaltyRegistry::new();
registry.push(AnalyticPenaltyKind::Isometry(Arc::new(
IsometryPenalty::new_euclidean(PsiSlice::full(n, Some(1)), 1),
)));
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![1.0_f64.ln()]]);
let sys = term
.assemble_arrow_schur(target.view(), &rho, Some(®istry))
.expect("assemble with isometry succeeds");
let bare = term
.assemble_arrow_schur(target.view(), &rho, None)
.expect("bare assemble succeeds");
let mut htt_iso = 0.0_f64;
for (r, b) in sys.rows.iter().zip(bare.rows.iter()) {
for (v, bv) in r.htt.iter().zip(b.htt.iter()) {
htt_iso += (v - bv) * (v - bv);
}
}
htt_iso.sqrt()
};
let base = isometry_curvature_norm(1.0);
assert!(
base > 1.0,
"the planted-circle fixture must produce a non-trivial isometry \
curvature block to make the scale test meaningful; got {base:.3e}"
);
for &lambda in &[3.0_f64, 10.0, 50.0] {
let scaled = isometry_curvature_norm(lambda);
let rel = (scaled - base).abs() / base;
assert!(
rel < 1.0e-6,
"SAE-assembled isometry curvature must be decoder-scale-invariant \
(#795): λ=1 → {base:.6e}, λ={lambda} → {scaled:.6e} (rel diff {rel:.3e}). \
A λ-dependent block is the un-normalized ‖B‖⁴ Gauss-Newton curvature \
whose mismatch with the scale-free gradient saturates the proximal \
ridge at 1e15."
);
}
}
#[test]
fn sae_isometry_joint_fit_is_physical_coscale_invariant_2099() {
use gam_problem::RowMetric;
use gam_terms::analytic_penalties::{
AnalyticPenaltyKind, AnalyticPenaltyRegistry, IsometryPenalty, PsiSlice,
};
let n = 24usize;
let p = 4usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.5) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let m = phi.ncols();
let base_decoder = Array2::from_shape_fn((m, p), |(b, c)| {
let scale = 1.0 / (1.0 + b as f64);
scale * ((b as f64 + 1.0) * (c as f64 + 1.0)).cos()
});
struct ScaleFit {
normalized_reconstruction: Array2<f64>,
criterion: f64,
components: [f64; 7],
}
let fit_at_scale = |physical_scale: f64| -> ScaleFit {
let scale_sq = physical_scale * physical_scale;
let decoder = &base_decoder * physical_scale;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"iso_converge",
SaeAtomBasisKind::Periodic,
1,
phi.clone(),
jet.clone(),
decoder.clone(),
Array2::<f64>::eye(m),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let target = phi.dot(&decoder);
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),
)
.expect("fixture assignment: one logit column and one coord block per atom");
let mut term = SaeManifoldTerm::new(vec![atom], assignment)
.expect("fixture term: every atom's basis width matches its assignment block");
let metric_factor = Array2::from_shape_fn((n, p * p), |(_, flat)| {
let output = flat / p;
let probe = flat % p;
if output == probe {
physical_scale.recip()
} else {
0.0
}
});
term.set_row_metric(
RowMetric::behavioral_fisher(Arc::new(metric_factor), p, p)
.expect("physical-unit precision metric is valid"),
)
.expect("physical-unit precision metric matches the output dimension");
let mut registry = AnalyticPenaltyRegistry::new();
registry.push(AnalyticPenaltyKind::Isometry(Arc::new(
IsometryPenalty::new_euclidean(PsiSlice::full(n, Some(1)), 1),
)));
let mut rho = SaeManifoldRho::new(0.0, (0.8_f64 / scale_sq).ln(), vec![array![0.0]]);
let loss = term
.run_joint_fit_arrow_schur(
target.view(),
&mut rho,
Some(®istry),
12,
1.0,
1.0e-4,
1.0e-4 / scale_sq,
)
.expect("joint fit with isometry gauge ON must converge at every decoder scale");
assert!(
loss.total().is_finite(),
"converged loss must be finite at c={physical_scale}, got {}",
loss.total()
);
let recon = term
.try_fitted_for_rho(&rho)
.expect("fitted reconstruction exists");
let criterion = term
.penalized_objective_total(target.view(), &rho, Some(®istry), 1.0)
.expect("co-scaled penalized criterion is defined");
assert!(
criterion.is_finite(),
"penalized criterion must be finite at c={physical_scale}, got {criterion}"
);
let scored_loss = term
.loss_scaled(target.view(), &rho, 1.0)
.expect("co-scaled loss breakdown is defined");
let analytic = term
.analytic_penalty_value_total(®istry, 1.0)
.expect("co-scaled analytic-penalty value is defined");
let repulsion = term.decoder_repulsion_value(1.0);
let separation = term.separation_barrier_value(1.0);
ScaleFit {
normalized_reconstruction: recon.mapv(|value| value / physical_scale),
criterion,
components: [
scored_loss.data_fit,
scored_loss.assignment_sparsity,
scored_loss.smoothness,
scored_loss.ard,
analytic,
repulsion,
separation,
],
}
};
let base = fit_at_scale(1.0);
let base_image_norm_sq = base
.normalized_reconstruction
.iter()
.map(|value| value * value)
.sum::<f64>();
assert!(
base_image_norm_sq.is_finite() && base_image_norm_sq > 0.0,
"unit-scale fitted image must be finite and nonzero"
);
for &physical_scale in &[5.0_f64, 25.0] {
let scaled = fit_at_scale(physical_scale);
let image_defect = base
.normalized_reconstruction
.iter()
.zip(scaled.normalized_reconstruction.iter())
.map(|(unit, rescaled)| {
let delta = unit - rescaled;
delta * delta
})
.sum::<f64>()
.sqrt()
/ base_image_norm_sq.sqrt();
let criterion_defect = (scaled.criterion - base.criterion).abs()
/ (1.0 + scaled.criterion.abs().max(base.criterion.abs()));
eprintln!(
"[#2099 fit co-scale] c={physical_scale}: image_defect={image_defect:.3e} \
criterion_defect={criterion_defect:.3e}; components \
[data,assignment,smooth,ard,analytic,repulsion,separation] base={:?} scaled={:?}",
base.components, scaled.components,
);
assert!(
image_defect < 1.0e-3,
"normalized fitted reconstruction changed under physical co-scale c={physical_scale}: \
relative image defect {image_defect:.3e}"
);
assert!(
criterion_defect < 1.0e-3,
"penalized criterion changed under physical co-scale c={physical_scale}: \
relative criterion defect {criterion_defect:.3e}"
);
let component_names = [
"data",
"assignment",
"smooth",
"ard",
"analytic",
"repulsion",
"separation",
];
for (idx, name) in component_names.into_iter().enumerate() {
let unit = base.components[idx];
let rescaled = scaled.components[idx];
let defect = (rescaled - unit).abs() / (1.0 + rescaled.abs().max(unit.abs()));
assert!(
defect < 1.0e-3,
"{name} criterion component changed under physical co-scale \
c={physical_scale}: unit={unit:.8e}, rescaled={rescaled:.8e}, \
relative defect={defect:.3e}"
);
}
}
}
#[test]
fn sae_single_planted_circle_embedded_isometry_fit_converges_795() {
use super::tests::{
PlantedCircleAssignmentMode, planted_circle_embedded, planted_circle_seed_term,
};
use gam_terms::analytic_penalties::{
AnalyticPenaltyKind, AnalyticPenaltyRegistry, IsometryPenalty, PsiSlice,
};
let n = 200usize;
let d_embed = 12usize;
let sigma = 0.02_f64;
let z = planted_circle_embedded(n, d_embed, sigma);
let (mut term, _seed_dispersion) =
planted_circle_seed_term(z.view(), PlantedCircleAssignmentMode::Softmax);
let mut registry = AnalyticPenaltyRegistry::new();
registry.push(AnalyticPenaltyKind::Isometry(Arc::new(
IsometryPenalty::new_euclidean(PsiSlice::full(n, Some(1)), 1),
)));
let mut rho = SaeManifoldRho::new(0.02_f64.ln(), 1.0_f64.ln(), vec![array![0.0_f64]]);
let loss = term
.run_joint_fit_arrow_schur(
z.view(),
&mut rho,
Some(®istry),
25,
0.04,
1.0e-6,
1.0e-6,
)
.expect(
"single planted circle embedded in D=12 with the isometry gauge ON must converge \
through the arrow-Schur joint fit (issue #795: the proximal ridge must not \
saturate at 1e15 / reject every step)",
);
assert!(
loss.total().is_finite(),
"#795: converged loss on the embedded planted circle must be finite, got {}",
loss.total()
);
let fitted = term.fitted();
assert!(
fitted.iter().all(|v| v.is_finite()),
"#795: fitted reconstruction must be finite"
);
let mut num = 0.0_f64;
let mut den = 0.0_f64;
for (r, t) in fitted.iter().zip(z.iter()) {
num += (r - t) * (r - t);
den += t * t;
}
let rel_recon = (num / den.max(1.0e-300)).sqrt();
assert!(
rel_recon < 0.4,
"#795: the isometry-gauged joint fit must recover the embedded planted circle \
(rel recon {rel_recon:.3e}); a residual near ‖z‖ is the ridge-saturation stall symptom"
);
}
#[test]
fn sae_k1_circle_penalized_quasi_laplace_criterion_ranks_fixed_point_2226() {
use super::tests::{
PlantedCircleAssignmentMode, planted_circle_embedded, planted_circle_seed_term,
};
use gam_terms::analytic_penalties::{
AnalyticPenaltyKind, AnalyticPenaltyRegistry, IsometryPenalty, PsiSlice,
};
let n = 200usize;
let d_embed = 12usize;
let sigma = 0.02_f64;
let z = planted_circle_embedded(n, d_embed, sigma);
let (mut term, _seed_dispersion) =
planted_circle_seed_term(z.view(), PlantedCircleAssignmentMode::Softmax);
let mut registry = AnalyticPenaltyRegistry::new();
registry.push(AnalyticPenaltyKind::Isometry(Arc::new(
IsometryPenalty::new_euclidean(PsiSlice::full(n, Some(1)), 1),
)));
let rho = SaeManifoldRho::new(0.02_f64.ln(), 1.0_f64.ln(), vec![array![0.0_f64]]);
let (value, loss, _cache) = term
.penalized_quasi_laplace_criterion_with_cache(
z.view(),
&rho,
Some(®istry),
200,
0.04,
1.0e-6,
1.0e-6,
)
.expect(
"#2226: the K=1 planted-circle inner solve reaches a numerical fixed point; \
penalized_quasi_laplace_criterion must rank that stationary iterate (affine-invariant Newton \
decrement) instead of refusing on an unreachable absolute gradient tolerance",
);
assert!(
value.is_finite() && loss.total().is_finite(),
"#2226: the ranked Laplace criterion must be finite (value={value}, loss={})",
loss.total()
);
}
#[test]
fn ranking_and_gradient_lanes_match_bare_reml() {
let mut objective = warmstart_test_objective_with_evaluator();
let rho_flat = objective.current_rho.to_flat();
let value_lane = objective
.eval_cost(&rho_flat)
.expect("value-probe lane evaluates penalized quasi-Laplace");
let mut objective_grad = warmstart_test_objective_with_evaluator();
let gradient_lane = objective_grad
.eval(&rho_flat)
.expect("gradient lane evaluates")
.cost;
assert!(
value_lane.is_finite() && gradient_lane.is_finite(),
"both lanes must be finite: value={value_lane}, gradient={gradient_lane}"
);
let bare_shared = {
let mut selected_term = objective_grad.term.clone();
let target = objective_grad.target.clone();
let rho_state = objective_grad
.baseline_rho
.from_flat(rho_flat.view())
.expect("the flattened rho has the length the baseline rho declares");
let (reml, _loss, _cache) = selected_term
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho_state,
objective_grad.registry.as_ref(),
objective_grad.inner_max_iter,
objective_grad.learning_rate,
objective_grad.ridge_ext_coord,
objective_grad.ridge_beta,
)
.expect("selected envelope state re-prices through the bare criterion");
reml
};
let value_vs_bare = (value_lane - bare_shared).abs();
assert!(
value_vs_bare < 1.0e-9,
"the ranking lane must report bare REML so selection and descent share \
one criterion: value_lane={value_lane}, bare={bare_shared}, \
diff={value_vs_bare}"
);
let gradient_vs_bare = (gradient_lane - bare_shared).abs();
assert!(
gradient_vs_bare < 1.0e-9,
"the gradient lane must report bare REML (no consistency fold), so its \
(cost, ∇f) pair is self-consistent for BFGS Armijo: \
gradient_lane={gradient_lane}, bare={bare_shared}, \
diff={gradient_vs_bare}"
);
assert!(
(value_lane - gradient_lane).abs() < 1.0e-9,
"ranking and gradient lanes must price one fixed point: \
value={value_lane}, gradient={gradient_lane}"
);
}
#[test]
fn amortized_warm_start_matches_or_beats_cold_inner_solve_on_known_manifold() {
let n = 24usize;
let p = 4usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.5) / n as f64);
let (phi, jet) = periodic_basis(&coords);
let m = phi.ncols();
let decoder = Array2::from_shape_fn((m, p), |(b, c)| {
let scale = 1.0 / (1.0 + b as f64);
scale * ((b as f64 + 1.0) * (c as f64 + 1.0)).cos()
});
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic_truth",
SaeAtomBasisKind::Periodic,
1,
phi.clone(),
jet,
decoder.clone(),
Array2::<f64>::eye(m),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let target = phi.dot(&decoder);
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("fixture assignment: one logit column and one coord block per atom");
let mut term = SaeManifoldTerm::new(vec![atom], assignment)
.expect("fixture term: every atom's basis width matches its assignment block");
let rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![1.0_f64.ln()]]);
let mut rho_cold = rho.clone();
term.run_joint_fit_arrow_schur(target.view(), &mut rho_cold, None, 12, 0.1, 1.0e-4, 1.0e-4)
.expect("cold inner solve converges on the known periodic manifold");
let cold_ev = {
let fitted = term
.try_fitted_for_rho(&rho_cold)
.expect("the fit above converged, so a fitted surface exists");
reconstruction_explained_variance(target.view(), fitted.view())
.expect("explained variance is defined for the planted target")
};
assert!(
cold_ev > 0.9,
"cold fit must recover the planted periodic manifold (EV={cold_ev})"
);
let warm_started = term
.warm_start_latents_from_amortized_encoder(target.view(), &rho_cold)
.expect("amortized warm-start runs on the fitted dictionary");
eprintln!("#1154 WARM-START: certified warm-started rows={warm_started}/{n}");
assert!(
warm_started <= n,
"the amortized encoder cannot warm-start more rows than the fitted \
batch size; warm_started={warm_started}, n={n}"
);
let mut rho_warm = rho.clone();
term.run_joint_fit_arrow_schur(target.view(), &mut rho_warm, None, 12, 0.1, 1.0e-4, 1.0e-4)
.expect("warm-started inner solve converges");
let warm_ev = {
let fitted = term
.try_fitted_for_rho(&rho_warm)
.expect("the fit above converged, so a fitted surface exists");
reconstruction_explained_variance(target.view(), fitted.view())
.expect("explained variance is defined for the planted target")
};
assert!(
warm_ev > 0.9,
"warm-started inner solve must still recover the planted manifold (warm_ev={warm_ev})"
);
assert!(
warm_ev >= cold_ev - 5.0e-3,
"amortized warm-start (co-trained inner solve) must recover the manifold \
about as well as the cold/sequential solve, to solver tolerance: \
warm_ev={warm_ev}, cold_ev={cold_ev}"
);
}
#[test]
fn sae_1026_curved_beats_linear_reconstruction_through_solver() {
let n = 48usize;
let p = 4usize;
let coords = Array2::from_shape_fn((n, 1), |(row, _)| (row as f64 + 0.5) / n as f64);
let (phi_c, jet_c) = periodic_basis(&coords);
let mc = phi_c.ncols();
let decoder_c = Array2::from_shape_fn((mc, p), |(b, c)| {
(1.0 / (1.0 + b as f64)) * ((b as f64 + 1.0) * (c as f64 + 1.0)).cos()
});
let target = phi_c.dot(&decoder_c);
let curved_ev = {
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"circle",
SaeAtomBasisKind::Periodic,
1,
phi_c.clone(),
jet_c,
decoder_c.clone(),
Array2::<f64>::eye(mc),
)
.expect("fixture atom: basis width, latent dim and decoder shape 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.clone()],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.expect("fixture assignment: one logit column and one coord block per atom");
let mut term = SaeManifoldTerm::new(vec![atom], assignment)
.expect("fixture term: every atom's basis width matches its assignment block");
let mut rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![array![1.0_f64.ln()]]);
term.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 12, 0.1, 1.0e-4, 1.0e-4)
.expect("curved inner solve converges on the planted circle");
let fitted = term
.try_fitted_for_rho(&rho)
.expect("the fit above converged, so a fitted surface exists");
reconstruction_explained_variance(target.view(), fitted.view())
.expect("target and reconstruction share a shape, so explained variance is defined")
};
let linear_ev = {
let evaluator = Arc::new(
EuclideanPatchEvaluator::new(1, 1)
.expect("a 1-D, 1-patch Euclidean basis is a valid evaluator"),
);
let (phi_l, jet_l) = evaluator
.evaluate(coords.view())
.expect("fixture coords are already wrapped into the evaluator's unit period");
let ml = phi_l.ncols();
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"linear",
SaeAtomBasisKind::EuclideanPatch,
1,
phi_l,
jet_l,
Array2::<f64>::zeros((ml, p)),
Array2::<f64>::eye(ml),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction")
.with_basis_second_jet(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords.clone()],
vec![LatentManifold::Euclidean],
AssignmentMode::softmax(1.0),
)
.expect("fixture assignment: one logit column and one coord block per atom");
let mut term = SaeManifoldTerm::new(vec![atom], assignment)
.expect("fixture term: every atom's basis width matches its assignment block");
let mut rho = SaeManifoldRho::new(0.0, 0.8_f64.ln(), vec![Array1::<f64>::zeros(1)]);
term.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 12, 0.1, 1.0e-4, 1.0e-4)
.expect("linear inner solve converges");
let fitted = term
.try_fitted_for_rho(&rho)
.expect("the fit above converged, so a fitted surface exists");
reconstruction_explained_variance(target.view(), fitted.view())
.expect("target and reconstruction share a shape, so explained variance is defined")
};
eprintln!("#1026 solver reconstruction: curved EV={curved_ev:.4}, linear EV={linear_ev:.4}");
assert!(
curved_ev > 0.9,
"the periodic atom must recover the planted circle through the solver (EV={curved_ev})"
);
assert!(
curved_ev > linear_ev + 0.2,
"curved must beat the matched-K linear baseline by a wide margin (the shatter \
penalty: a degree-1 secant cannot follow a closed circle): \
curved={curved_ev}, linear={linear_ev}"
);
}
#[test]
fn sae_1026_solver_recovers_separable_superposition_but_not_below_2k() {
let recover = |p: usize, overlap: bool| -> f64 {
let n = 80usize;
let theta_a = Array2::from_shape_fn((n, 1), |(r, _)| ((r as f64) * 0.043).rem_euclid(1.0));
let theta_b =
Array2::from_shape_fn((n, 1), |(r, _)| ((r as f64) * 0.071 + 0.13).rem_euclid(1.0));
let mut target = Array2::<f64>::zeros((n, p));
for r in 0..n {
let a = std::f64::consts::TAU * theta_a[[r, 0]];
let b = std::f64::consts::TAU * theta_b[[r, 0]];
if !overlap {
target[[r, 0]] = a.cos();
target[[r, 1]] = a.sin();
target[[r, 2]] = b.cos();
target[[r, 3]] = b.sin();
} else {
target[[r, 0]] += a.cos();
target[[r, 1]] += a.sin();
target[[r, 1]] += b.cos();
target[[r, 2]] += b.sin();
}
}
let seed_a =
Array2::from_shape_fn((n, 1), |(r, _)| (theta_a[[r, 0]] + 0.03).rem_euclid(1.0));
let seed_b =
Array2::from_shape_fn((n, 1), |(r, _)| (theta_b[[r, 0]] + 0.03).rem_euclid(1.0));
let (pa, ja) = periodic_basis(&seed_a);
let (pb, jb) = periodic_basis(&seed_b);
let m = pa.ncols();
let a0 = SaeManifoldAtom::new_with_provided_function_gram(
"cA",
SaeAtomBasisKind::Periodic,
1,
pa,
ja,
Array2::<f64>::zeros((m, p)),
Array2::<f64>::eye(m),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let a1 = SaeManifoldAtom::new_with_provided_function_gram(
"cB",
SaeAtomBasisKind::Periodic,
1,
pb,
jb,
Array2::<f64>::zeros((m, p)),
Array2::<f64>::eye(m),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction")
.with_basis_evaluator(Arc::new(TestPeriodicEvaluator));
let logits = Array2::<f64>::from_elem((n, 2), 6.0 * 0.5);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![seed_a.clone(), seed_b.clone()],
vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
],
AssignmentMode::ordered_beta_bernoulli(0.5, 1.0, false),
)
.expect("fixture assignment: one logit column and one coord block per atom");
let mut term = SaeManifoldTerm::new(vec![a0, a1], assignment)
.expect("fixture term: every atom's basis width matches its assignment block");
let mut rho = SaeManifoldRho::new(
0.0,
0.01_f64.ln(),
vec![array![1.0_f64.ln()], array![1.0_f64.ln()]],
);
term.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 24, 0.1, 1.0e-4, 1.0e-4)
.expect("K=2 inner solve converges");
let fitted = term
.try_fitted_for_rho(&rho)
.expect("the fit above converged, so a fitted surface exists");
reconstruction_explained_variance(target.view(), fitted.view())
.expect("target and reconstruction share a shape, so explained variance is defined")
};
let separable = recover(4, false);
let under_determined = recover(3, true);
eprintln!(
"#1026 K=2 superposition: separable(p=4)={separable:.4}, overlap(p=3)={under_determined:.4}"
);
assert!(
separable > 0.95,
"the joint solver must recover two superposed circles when p >= 2K (EV={separable})"
);
assert!(
under_determined > 0.9,
"amplitude-aware routing reconstructs even the p < 2K overlapping \
superposition (reconstruction, not identifiability): overlap EV={under_determined}"
);
}
#[test]
pub(crate) fn deflated_solver_matches_plain_solve_when_no_gauge_is_installed() {
let cache = diagonal_latent_cache(&[2.0_f64, 5.0, 7.0]);
let solver = DeflatedArrowSolver::plain(&cache);
let rhs_t = array![4.0_f64, 10.0, -14.0];
let rhs_beta = Array1::<f64>::zeros(0);
let (plain_t, plain_beta) = cache
.full_inverse_apply(rhs_t.view(), rhs_beta.view())
.expect("plain cache solve");
let solved = solver
.solve(rhs_t.view(), rhs_beta.view())
.expect("adapter solve");
assert_eq!(solved.t.len(), plain_t.len());
for idx in 0..plain_t.len() {
assert_abs_diff_eq!(solved.t[idx], plain_t[idx], epsilon = 0.0);
}
assert_eq!(solved.beta.len(), plain_beta.len());
for idx in 0..plain_beta.len() {
assert_abs_diff_eq!(solved.beta[idx], plain_beta[idx], epsilon = 0.0);
}
}
#[test]
pub(crate) fn deflated_solver_matches_dense_quotient_pseudoinverse_on_near_null_fixture() {
let cache = diagonal_latent_cache(&[2.0_f64, 1.0e-14]);
let gauge = array![0.0_f64, 1.0];
let solver = DeflatedArrowSolver::from_orthonormal_gauges(&cache, vec![gauge], 2.0)
.expect("deflated solver");
let rhs_beta = Array1::<f64>::zeros(0);
let physical_rhs = array![4.0_f64, 0.0];
let solved = solver
.solve(physical_rhs.view(), rhs_beta.view())
.expect("physical solve");
let oracle = array![2.0_f64, 0.0];
for idx in 0..oracle.len() {
assert_abs_diff_eq!(solved.t[idx], oracle[idx], epsilon = 1.0e-12);
}
let gauge_rhs = array![0.0_f64, 1.0];
let plain = cache
.full_inverse_apply(gauge_rhs.view(), rhs_beta.view())
.expect("plain gauge solve")
.0;
let stiffened = solver
.solve(gauge_rhs.view(), rhs_beta.view())
.expect("stiffened gauge solve")
.t;
assert!(plain[1] > 1.0e13, "plain near-null solve must be huge");
assert_abs_diff_eq!(stiffened[1], 0.5, epsilon = 1.0e-12);
}
#[test]
pub(crate) fn gauge_fixed_krylov_operator_matches_deflated_preconditioner_2253() {
let cache = diagonal_latent_cache(&[2.0_f64, 1.0e-14]);
let gauge = array![0.0_f64, 1.0];
let stiffness = 2.0;
let solver =
DeflatedArrowSolver::from_orthonormal_gauges(&cache, vec![gauge.clone()], stiffness)
.expect("deflated solver");
let rhs = SaeArrowVector {
t: array![4.0_f64, 1.0],
beta: Array1::zeros(0),
};
let raw_a = |v: &SaeArrowVector| -> Result<SaeArrowVector, String> {
Ok(SaeArrowVector {
t: array![2.0 * v.t[0], 0.0],
beta: Array1::zeros(0),
})
};
let add_gauge_stiffness = |v: &SaeArrowVector, applied: &mut SaeArrowVector| {
let coefficient = stiffness * gauge.dot(&v.t);
for index in 0..gauge.len() {
applied.t[index] += coefficient * gauge[index];
}
};
let raw = solve_b_preconditioned_gmres_with(
&rhs,
|v| raw_a(v),
|vector| solver.solve(vector.t.view(), vector.beta.view()),
);
assert!(
raw.is_err(),
"raw A with rhs mass on its exact gauge null must not pass the residual certificate"
);
let solved = solve_b_preconditioned_gmres_with(
&rhs,
|v| {
let mut out = raw_a(v)?;
add_gauge_stiffness(v, &mut out);
Ok(out)
},
|vector| solver.solve(vector.t.view(), vector.beta.view()),
)
.expect("gauge-fixed exact-stationarity solve");
assert_abs_diff_eq!(solved.t[0], 2.0, epsilon = 1.0e-12);
assert_abs_diff_eq!(solved.t[1], 0.5, epsilon = 1.0e-12);
let mut applied = raw_a(&solved).expect("raw A apply");
add_gauge_stiffness(&solved, &mut applied);
let residual = SaeArrowVector {
t: &applied.t - &rhs.t,
beta: &applied.beta - &rhs.beta,
};
assert!(
sae_norm(&residual) <= 1.0e-12 * sae_norm(&rhs).max(1.0),
"gauge-fixed operator and inverse must satisfy the original residual; got {:.3e}",
sae_norm(&residual)
);
}
#[test]
pub(crate) fn pca_seed_handles_huge_equal_finite_columns_without_mean_overflow() {
let z = array![[1.0e308_f64, 1.0e308], [1.0e308, 1.0e308]];
let coords = sae_pca_seed_initial_coords(z.view(), &[SaeAtomBasisKind::Periodic], &[1])
.expect("the planted target has full enough rank to seed one 1-D atom");
assert_eq!(coords.dim(), (1, 2, 1));
assert!(
coords.iter().all(|value| value.is_finite()),
"huge finite equal columns must not overflow the PCA seed mean: {coords:?}"
);
}
#[test]
pub(crate) fn pca_seed_rejects_huge_finite_span_that_overflows_centering() {
let z = array![[1.0e308_f64, 0.0], [-1.0e308, 0.0]];
let err = sae_pca_seed_initial_coords(z.view(), &[SaeAtomBasisKind::Periodic], &[1])
.expect_err("opposite huge finite values exceed f64 centering range");
assert!(
err.contains("centered Z is non-finite") || err.contains("SVD failed"),
"unexpected PCA seed error: {err}"
);
}
#[test]
fn smooth_threshold_hdiag_third_derivative_matches_central_difference_1415() {
use ndarray::{Array1, Array2, Array3};
let n = 6usize;
let k = 2usize;
let p = 3usize;
let temperature = 0.35_f64;
let threshold = 0.1_f64;
let logits = Array2::<f64>::from_shape_vec(
(n, k),
vec![
0.1, 0.0, 0.2, -0.05, 0.05, 0.15, 0.25, 0.3, -0.1, 0.12, 0.18, 0.08,
],
)
.expect("valid logit grid");
let atoms: Vec<SaeManifoldAtom> = (0..k)
.map(|i| {
SaeManifoldAtom::new_with_provided_function_gram(
&format!("atom{i}"),
SaeAtomBasisKind::EuclideanPatch,
1,
Array2::<f64>::ones((n, 2)),
Array3::<f64>::zeros((n, 2, 1)),
Array2::<f64>::zeros((2, p)),
Array2::<f64>::eye(2),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction")
})
.collect();
let coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
let manifolds = vec![LatentManifold::Euclidean; k];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits.clone(),
coords,
manifolds,
AssignmentMode::threshold_gate(temperature, threshold),
)
.expect("valid smooth threshold assignment");
let term = SaeManifoldTerm::new(atoms, assignment)
.expect("fixture term: every atom's basis width matches its assignment block");
let rho = SaeManifoldRho::new(0.7_f64.ln(), -6.0, vec![Array1::<f64>::zeros(1); k]);
let inv_tau = 1.0 / temperature;
let sparsity = rho
.lambda_sparse()
.expect("the fixture rho was built with a sparsity lambda");
let p2 = |logit: f64| -> f64 {
let a = gam_linalg::utils::stable_logistic((logit - threshold) * inv_tau);
let s = a * (1.0 - a);
sparsity * s * (1.0 - 2.0 * a) * inv_tau * inv_tau
};
let mut saw_threshold = false;
for row in 0..n {
for atom in 0..k {
let logit = logits[[row, atom]];
let entry = term.assignment_prior_hdiag_derivative_entry(
sparsity,
row,
atom,
SaeLocalRowVar::Logit { atom },
None,
);
let h = 1.0e-3_f64;
let fd = (-p2(logit + 2.0 * h) + 8.0 * p2(logit + h) - 8.0 * p2(logit - h)
+ p2(logit - 2.0 * h))
/ (12.0 * h);
let scale = entry.abs().max(fd.abs()).max(1.0e-8);
assert!(
(entry - fd).abs() <= 1.0e-5 * scale,
"row {row} atom {atom}: P''' entry {entry:e} vs FD {fd:e}"
);
if (logit - threshold).abs() < 1e-12 {
saw_threshold = true;
let expected = -sparsity / 8.0 * inv_tau * inv_tau * inv_tau;
assert_abs_diff_eq!(entry, expected, epsilon = 1e-9);
assert!(
entry < -1e-6,
"threshold third derivative must be strictly negative (old buggy \
formula returned 0): entry={entry:e}"
);
}
}
}
assert!(
saw_threshold,
"fixture must include a logit exactly at the threshold to pin −λ/(8τ³)"
);
}
#[test]
fn encode_grad_hess_and_beta_eta_match_finite_differences() {
use crate::encode::{beta_eta_newton, encode_grad_hess};
use ndarray::Array2;
let train = Array2::from_shape_fn((24, 1), |(r, _)| (r as f64 + 0.5) / 24.0);
let (phi, jet) = periodic_basis(&train);
let m = phi.ncols();
let p = 4usize;
let decoder = Array2::from_shape_fn((m, p), |(b, c)| {
(1.0 / (1.0 + b as f64)) * ((b as f64 + 1.0) * (c as f64 + 1.0)).cos()
});
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"circle",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder.clone(),
Array2::<f64>::eye(m),
)
.expect("fixture atom: basis width, latent dim and decoder shape agree by construction");
let eval = TestPeriodicEvaluator;
let amplitude = 0.8_f64;
let decode = |t: f64| -> ndarray::Array1<f64> {
let coords = Array2::from_shape_fn((1, 1), |_| t);
let (ph, _) = periodic_basis(&coords);
amplitude * ph.dot(&decoder).row(0).to_owned()
};
let t0 = 0.137_f64;
let x = &decode(0.42) + &ndarray::Array1::from_vec(vec![0.3, -0.2, 0.15, -0.25]);
let f = |t: f64| -> f64 {
let r = &decode(t) - &x;
0.5 * r.dot(&r)
};
let t_view = ndarray::Array1::from_vec(vec![t0]);
let (g, h) = encode_grad_hess(&atom, &eval, t_view.view(), x.view(), amplitude)
.expect("encode_grad_hess runs")
.expect("second jet present ⇒ Some");
let eps = 1e-6;
let g_fd = (f(t0 + eps) - f(t0 - eps)) / (2.0 * eps);
assert_abs_diff_eq!(g[0], g_fd, epsilon = 1e-6);
let h_fd = (f(t0 + eps) - 2.0 * f(t0) + f(t0 - eps)) / (eps * eps);
assert_abs_diff_eq!(h[[0, 0]], h_fd, epsilon = 5e-3);
let mut hpd = h.clone();
if hpd[[0, 0]] <= 0.0 {
hpd[[0, 0]] = 1.5;
}
let (beta, eta, delta) = beta_eta_newton(hpd.view(), g.view())
.expect("beta_eta_newton runs")
.expect("SPD ⇒ Some");
assert_abs_diff_eq!(beta * hpd[[0, 0]], 1.0, epsilon = 1e-12);
assert_abs_diff_eq!(delta[0], -g[0] / hpd[[0, 0]], epsilon = 1e-12);
assert_abs_diff_eq!(eta, (g[0] / hpd[[0, 0]]).abs(), epsilon = 1e-12);
}
#[test]
pub(crate) fn run_joint_fit_max_iter_zero_freezes_beta_verbatim() {
let mut term = warmstart_test_objective().term;
let dim = term.beta_dim();
let pristine = term.flatten_beta();
let seed: Array1<f64> = Array1::from_shape_fn(dim, |i| pristine[i] + 0.5 + 0.01 * (i as f64));
assert!(
(&seed - &pristine).iter().any(|d| d.abs() > 1e-6),
"seed must differ from the pristine β for the freeze check to be meaningful"
);
term.set_flat_beta(seed.view())
.expect("length-matching β must install");
let target = array![[0.20_f64], [-0.10], [0.30], [0.05]];
let mut rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1)]);
let loss = term
.run_joint_fit_arrow_schur(target.view(), &mut rho, None, 0, 1.0, 1.0e-6, 1.0e-6)
.expect("zero-iteration joint fit at the warm-started β must succeed");
assert!(
loss.total().is_finite(),
"frozen-state loss must be finite; got {}",
loss.total()
);
let frozen = term.flatten_beta();
for (i, (&f, &s)) in frozen.iter().zip(seed.iter()).enumerate() {
assert!(
(f - s).abs() < 1e-12,
"max_iter==0 must freeze β verbatim at coord {i}: frozen {f} != seed {s} (#850)"
);
}
}