use super::derivative_oracle::{
BranchCertificate, EigenDerivativeRoute, MajorizerAnchorMode, PivotBranch,
};
use super::dual::DualKinkBranch;
use super::*;
use approx::assert_abs_diff_eq;
use gam_solve::arrow_schur::{
ArrowFactorSlab, ArrowHtbetaCache, ArrowPcgDiagnostics, ArrowSolverMode, ArrowUndampedFactors,
};
use gam_solve::evidence::arrow_log_det_from_cache;
use ndarray::array;
#[test]
pub(crate) fn sae_torus_atom_recovers_two_frequency_synthetic() {
let n = 96usize;
let p = 4usize;
let h = 3usize;
let d = 2usize;
let evaluator = TorusHarmonicEvaluator::new(d, h).unwrap();
let m = evaluator.basis_size();
let mut true_coords = Array2::<f64>::zeros((n, d));
for i in 0..n {
true_coords[[i, 0]] = ((i as f64) * 0.137).rem_euclid(1.0);
true_coords[[i, 1]] = ((i as f64) * 0.241 + 0.13).rem_euclid(1.0);
}
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let t1 = 2.0 * std::f64::consts::PI * true_coords[[i, 0]];
let t2 = 2.0 * std::f64::consts::PI * true_coords[[i, 1]];
z[[i, 0]] = t1.sin() + 0.3 * t2.cos();
z[[i, 1]] = t1.cos() + 0.2 * (t1 + t2).sin();
z[[i, 2]] = t2.sin();
z[[i, 3]] = 0.5 * (t1 - t2).cos();
}
let sst: f64 = z.iter().map(|v| v * v).sum::<f64>();
let (phi0, jet0) = evaluator.evaluate(true_coords.view()).unwrap();
let mut penalty = Array2::<f64>::eye(m);
penalty *= 1.0e-4;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"torus_atom",
SaeAtomBasisKind::Torus,
d,
phi0,
jet0,
Array2::<f64>::zeros((m, p)),
penalty,
)
.unwrap()
.with_basis_evaluator(Arc::new(TorusHarmonicEvaluator::new(d, h).unwrap()));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![true_coords],
vec![LatentManifold::Product(vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
])],
AssignmentMode::softmax(0.5),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let mut rho = SaeManifoldRho::new(0.0, -4.0, vec![Array1::<f64>::zeros(d)]);
let ridge = 1.0e-6;
for _ in 0..10 {
let loss = term
.run_joint_fit_arrow_schur(z.view(), &mut rho, None, 1, 1.0, ridge, ridge)
.unwrap();
if !loss.total().is_finite() {
break;
}
}
let fitted = term.fitted();
assert_eq!(fitted.dim(), (n, p));
let mut sse = 0.0_f64;
for ((row, col), v) in fitted.indexed_iter() {
let r = v - z[[row, col]];
sse += r * r;
}
let r2 = 1.0 - sse / sst.max(1.0e-12);
assert!(
r2 >= 0.5,
"torus atom R² too low: {r2:.4} (sst={sst:.4}, sse={sse:.4})"
);
}
#[test]
pub(crate) fn sae_sphere_atom_recovers_synthetic_signal() {
let n = 96usize;
let p = 3usize;
let d = 2usize;
let mut true_coords = Array2::<f64>::zeros((n, d));
for i in 0..n {
let t = (i as f64) / (n as f64);
true_coords[[i, 0]] = -0.5 + 1.0 * t; true_coords[[i, 1]] = -std::f64::consts::PI + 2.0 * std::f64::consts::PI * t;
}
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let lat = true_coords[[i, 0]];
let lon = true_coords[[i, 1]];
let x = lat.cos() * lon.cos();
let y = lat.cos() * lon.sin();
let zc = lat.sin();
z[[i, 0]] = x;
z[[i, 1]] = y;
z[[i, 2]] = zc;
}
let sst: f64 = z.iter().map(|v| v * v).sum::<f64>();
let (phi0, jet0) = SphereChartEvaluator.evaluate(true_coords.view()).unwrap();
let m = phi0.ncols();
let mut penalty = Array2::<f64>::eye(m);
penalty *= 1.0e-4;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"sphere_atom",
SaeAtomBasisKind::Sphere,
d,
phi0,
jet0,
Array2::<f64>::zeros((m, p)),
penalty,
)
.unwrap()
.with_basis_evaluator(Arc::new(SphereChartEvaluator));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![true_coords],
vec![LatentManifold::Product(vec![
LatentManifold::Interval {
lo: -std::f64::consts::FRAC_PI_2,
hi: std::f64::consts::FRAC_PI_2,
},
LatentManifold::Circle {
period: std::f64::consts::TAU,
},
])],
AssignmentMode::softmax(0.5),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let mut rho = SaeManifoldRho::new(0.0, -4.0, vec![Array1::<f64>::zeros(2)]);
let ridge = 1.0e-6;
for _ in 0..10 {
let loss = term
.run_joint_fit_arrow_schur(z.view(), &mut rho, None, 1, 1.0, ridge, ridge)
.unwrap();
if !loss.total().is_finite() {
break;
}
}
let fitted = term.fitted();
assert_eq!(fitted.dim(), (n, p));
let mut sse = 0.0_f64;
for ((row, col), v) in fitted.indexed_iter() {
let r = v - z[[row, col]];
sse += r * r;
}
let r2 = 1.0 - sse / sst.max(1.0e-12);
assert!(
r2 >= 0.5,
"sphere atom R² too low: {r2:.4} (sst={sst:.4}, sse={sse:.4})"
);
}
#[test]
pub(crate) fn sae_manifold_fit_10_steps_one_harmonic_reaches_high_r2() {
let n = 64usize;
let m = 3usize;
let p = 1usize;
let true_t: Vec<f64> = (0..n).map(|i| (i as f64) / (n as f64)).collect();
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let angle = 2.0 * std::f64::consts::PI * true_t[i];
z[[i, 0]] = 0.7 * angle.sin() + 0.3 * angle.cos();
}
let sst: f64 = z.iter().map(|v| v * v).sum::<f64>();
let evaluator = PeriodicHarmonicEvaluator::new(m).unwrap();
let mut coords0_data = Array2::<f64>::zeros((n, 1));
for i in 0..n {
coords0_data[[i, 0]] = (true_t[i] + 0.25).rem_euclid(1.0);
}
let (phi0, jet0) = evaluator.evaluate(coords0_data.view()).unwrap();
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic_atom",
SaeAtomBasisKind::Periodic,
1,
phi0,
jet0,
Array2::<f64>::zeros((m, p)),
Array2::<f64>::eye(m),
)
.unwrap()
.with_basis_evaluator(Arc::new(PeriodicHarmonicEvaluator::new(m).unwrap()));
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords0_data],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(0.5),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let mut rho = SaeManifoldRho::new(0.0, -6.0, vec![Array1::<f64>::zeros(1)]);
let max_iter = 10usize;
let learning_rate = 1.0;
let ridge = 1.0e-6;
let mut prev_total = f64::INFINITY;
for _ in 0..max_iter {
let loss = term
.run_joint_fit_arrow_schur(z.view(), &mut rho, None, 1, learning_rate, ridge, ridge)
.unwrap();
let total = loss.total();
if !total.is_finite() {
break;
}
let denom = prev_total.abs().max(1.0e-12);
let rel = (prev_total - total).abs() / denom;
prev_total = total;
if rel < 1.0e-6 {
break;
}
}
let fitted = term.fitted();
assert_eq!(fitted.dim(), (n, p));
let mut ssr = 0.0;
for i in 0..n {
let r = z[[i, 0]] - fitted[[i, 0]];
ssr += r * r;
}
let r2 = 1.0 - ssr / sst.max(1.0e-12);
assert!(
r2 >= 0.95,
"10-step in-sample R² = {r2:.4} (ssr={ssr:.6}, sst={sst:.6}) should be >= 0.95"
);
}
#[test]
pub(crate) fn softmax_assignment_hessian_diag_is_available_for_k2() {
let n = 4usize;
let k = 2usize;
let logits = Array2::<f64>::from_shape_fn((n, k), |(i, j)| 0.1 * (i as f64) - 0.2 * (j as f64));
let coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
let manifolds = vec![LatentManifold::Circle { period: 1.0 }; k];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coords,
manifolds,
AssignmentMode::softmax(0.7),
)
.unwrap();
let rho = SaeManifoldRho::new(0.0, -6.0, vec![Array1::<f64>::zeros(1); k]);
let (grad, diag) = assignment_prior_grad_hdiag(&assignment, &rho)
.expect("softmax assignment Hessian diagonal must be available");
assert_eq!(grad.len(), n * k);
assert_eq!(diag.len(), n * k);
assert!(grad.iter().all(|v| v.is_finite()));
assert!(diag.iter().all(|v| v.is_finite()));
}
#[test]
pub(crate) fn sae_registry_refuses_assignment_sparsity_penalties() {
let n = 3usize;
let k = 2usize;
let logits = Array2::<f64>::zeros((n, k));
let coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
let manifolds = vec![LatentManifold::Circle { period: 1.0 }; k];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coords,
manifolds,
AssignmentMode::softmax(0.7),
)
.expect("valid assignment");
let atoms: Vec<SaeManifoldAtom> = (0..k)
.map(|atom_idx| {
SaeManifoldAtom::new_with_provided_function_gram(
format!("periodic_{atom_idx}"),
SaeAtomBasisKind::Periodic,
1,
Array2::<f64>::ones((n, 1)),
Array3::<f64>::zeros((n, 1, 1)),
Array2::<f64>::zeros((1, 1)),
Array2::<f64>::eye(1),
)
.expect("valid atom")
})
.collect();
let term = SaeManifoldTerm::new(atoms, assignment).expect("valid SAE term");
let mut softmax_registry = AnalyticPenaltyRegistry::new();
softmax_registry.push(AnalyticPenaltyKind::SoftmaxAssignmentSparsity(Arc::new(
gam_terms::analytic_penalties::SoftmaxAssignmentSparsityPenalty::new(k, 0.7),
)));
let softmax_err = term
.validate_analytic_penalty_registry(&softmax_registry)
.expect_err("SAE registry must reject softmax assignment sparsity");
assert!(softmax_err.contains("assignment sparsity"));
let mut ordered_beta_bernoulli_registry = AnalyticPenaltyRegistry::new();
ordered_beta_bernoulli_registry.push(AnalyticPenaltyKind::OrderedBetaBernoulli(Arc::new(
gam_terms::analytic_penalties::OrderedBetaBernoulliPenalty::new(k, 1.2, 0.7, false),
)));
let ordered_beta_bernoulli_err = term
.validate_analytic_penalty_registry(&ordered_beta_bernoulli_registry)
.expect_err("SAE registry must reject ordered Beta--Bernoulli assignment sparsity");
assert!(ordered_beta_bernoulli_err.contains("assignment sparsity"));
}
#[test]
pub(crate) fn ordered_beta_bernoulli_fixed_alpha_assignment_value_matches_logit_gradient_fd() {
let n = 4usize;
let k = 3usize;
let logits = Array2::<f64>::from_shape_vec(
(n, k),
vec![
-0.4, 0.2, 0.7, 0.1, -0.3, 0.5, 0.8, -0.1, -0.6, 0.3, 0.6, -0.2,
],
)
.expect("valid ordered Beta--Bernoulli logit grid");
let coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
let manifolds = vec![LatentManifold::Circle { period: 1.0 }; k];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coords,
manifolds,
AssignmentMode::ordered_beta_bernoulli(0.9, 1.4, false),
)
.expect("valid ordered Beta--Bernoulli assignment");
let rho = SaeManifoldRho::new(0.23_f64.ln(), -6.0, vec![Array1::<f64>::zeros(1); k]);
let (grad, _) = assignment_prior_grad_hdiag(&assignment, &rho)
.expect("ordered Beta--Bernoulli assignment gradient");
let idx = 5usize;
let step = 1.0e-6_f64;
let mut plus = assignment.clone();
plus.logits[[idx / k, idx % k]] += step;
let mut minus = assignment.clone();
minus.logits[[idx / k, idx % k]] -= step;
let fd = (assignment_prior_value(&plus, &rho).unwrap()
- assignment_prior_value(&minus, &rho).unwrap())
/ (2.0 * step);
assert_abs_diff_eq!(grad[idx], fd, epsilon = 2.0e-7);
}
#[test]
pub(crate) fn threshold_gate_assignment_value_matches_logit_gradient_fd() {
let n = 4usize;
let k = 2usize;
let temperature = 0.35_f64;
let threshold = 0.1_f64;
let logits =
Array2::<f64>::from_shape_vec((n, k), vec![-13.0, -0.2, 0.0, 0.05, 0.15, 0.4, 0.9, 1.5])
.expect("valid ThresholdGate logit grid");
let coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
let manifolds = vec![LatentManifold::Circle { period: 1.0 }; k];
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coords,
manifolds,
AssignmentMode::threshold_gate(temperature, threshold),
)
.expect("valid ThresholdGate assignment");
let rho = SaeManifoldRho::new(0.7_f64.ln(), -6.0, vec![Array1::<f64>::zeros(1); k]);
let (grad, _) =
assignment_prior_grad_hdiag(&assignment, &rho).expect("ThresholdGate assignment gradient");
let idx = 4usize;
let step = 1.0e-6_f64;
let mut plus = assignment.clone();
plus.logits[[idx / k, idx % k]] += step;
let mut minus = assignment.clone();
minus.logits[[idx / k, idx % k]] -= step;
let fd = (assignment_prior_value(&plus, &rho).unwrap()
- assignment_prior_value(&minus, &rho).unwrap())
/ (2.0 * step);
assert_abs_diff_eq!(grad[idx], fd, epsilon = 2.0e-8);
}
#[test]
pub(crate) fn threshold_gate_assignment_prior_hessian_diag_is_exact_over_logit_sweep() {
let n = 6usize;
let k = 2usize;
let temperature = 0.35_f64;
let threshold = 0.1_f64;
let logits = Array2::<f64>::from_shape_vec(
(n, k),
vec![
-2.0, -0.2, 0.0, 0.05, 0.1, 0.15, 0.4, 0.9, 1.5, 2.5, 4.0, 6.0,
],
)
.expect("valid logit grid");
let coords: Vec<Array2<f64>> = (0..k).map(|_| Array2::<f64>::zeros((n, 1))).collect();
let manifolds = vec![LatentManifold::Circle { period: 1.0 }; 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 rho = SaeManifoldRho::new(0.7_f64.ln(), -6.0, vec![Array1::<f64>::zeros(1); k]);
let (grad, diag) = assignment_prior_grad_hdiag(&assignment, &rho)
.expect("smooth threshold assignment prior hessian diag");
let inv_tau = 1.0 / temperature;
let inv_tau2 = inv_tau * inv_tau;
let sparsity_strength = rho.lambda_sparse().unwrap();
assert_eq!(grad.len(), n * k);
assert_eq!(diag.len(), n * k);
let mut saw_negative = false;
for (idx, &entry) in diag.iter().enumerate() {
let logit = logits[[idx / k, idx % k]];
let activation = gam_linalg::utils::stable_logistic((logit - threshold) * inv_tau);
let slope = activation * (1.0 - activation);
let expected = sparsity_strength * slope * (1.0 - 2.0 * activation) * inv_tau2;
assert!(
entry.is_finite(),
"threshold-gate hessian_diag must be finite at index {idx}"
);
saw_negative |= entry < 0.0;
assert_abs_diff_eq!(entry, expected, epsilon = 1e-12);
}
assert!(
saw_negative,
"exact threshold-gate hessian_diag must go negative above the threshold"
);
}
#[test]
pub(crate) fn ordered_beta_bernoulli_k2_periodic_torus_recovers_signal_with_lsq_init() {
use faer::Side as FaerSide;
use gam_linalg::faer_ndarray::{FaerCholesky, fast_ata, fast_atb};
let n = 200usize;
let p = 8usize;
let k = 2usize;
let m = 5usize;
let mut theta = Array2::<f64>::zeros((n, 2));
for i in 0..n {
theta[[i, 0]] = ((i as f64) * 0.07) % 1.0;
theta[[i, 1]] = ((i as f64) * 0.13 + 0.31) % 1.0;
}
let mut raw = Array2::<f64>::zeros((n, 4));
for i in 0..n {
let a1 = 2.0 * std::f64::consts::PI * theta[[i, 0]];
let a2 = 2.0 * std::f64::consts::PI * theta[[i, 1]];
raw[[i, 0]] = a1.cos();
raw[[i, 1]] = a1.sin();
raw[[i, 2]] = a2.cos();
raw[[i, 3]] = a2.sin();
}
let mix = Array2::<f64>::from_shape_fn((4, p), |(i, j)| {
((i as f64 + 1.0) * 0.37 + (j as f64) * 0.21).sin()
});
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
for j in 0..p {
let mut acc = 0.0;
for r in 0..4 {
acc += raw[[i, r]] * mix[[r, j]];
}
z[[i, j]] = acc;
}
}
let mut col_mean = Array1::<f64>::zeros(p);
for j in 0..p {
let mut acc = 0.0;
for i in 0..n {
acc += z[[i, j]];
}
col_mean[j] = acc / n as f64;
}
for i in 0..n {
for j in 0..p {
z[[i, j]] -= col_mean[j];
}
}
let mut coords_k = vec![Array2::<f64>::zeros((n, 1)); k];
for i in 0..n {
coords_k[0][[i, 0]] = (theta[[i, 0]] + 0.05).rem_euclid(1.0);
coords_k[1][[i, 0]] = (theta[[i, 1]] + 0.07).rem_euclid(1.0);
}
let evaluator = PeriodicHarmonicEvaluator::new(m).unwrap();
let mut phi_k = Vec::with_capacity(k);
let mut jet_k = Vec::with_capacity(k);
for atom_idx in 0..k {
let (phi, jet) = evaluator.evaluate(coords_k[atom_idx].view()).unwrap();
phi_k.push(phi);
jet_k.push(jet);
}
let m_total = k * m;
let mut x = Array2::<f64>::zeros((n, m_total));
for atom_idx in 0..k {
for i in 0..n {
for col in 0..m {
x[[i, atom_idx * m + col]] = 0.5 * phi_k[atom_idx][[i, col]];
}
}
}
let mut xtx = fast_ata(&x);
let mut trace = 0.0_f64;
for i in 0..m_total {
trace += xtx[[i, i]];
}
let jitter = (trace / m_total as f64).max(1.0) * 1.0e-8;
for i in 0..m_total {
xtx[[i, i]] += jitter;
}
let xtz = fast_atb(&x, &z);
let b_joint = xtx
.cholesky(FaerSide::Lower)
.expect("LSQ Cholesky")
.solve_mat(&xtz);
let mut atoms = Vec::with_capacity(k);
for atom_idx in 0..k {
let mut b = Array2::<f64>::zeros((m, p));
for col in 0..m {
for j in 0..p {
b[[col, j]] = b_joint[[atom_idx * m + col, j]];
}
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
format!("torus_atom_{atom_idx}"),
SaeAtomBasisKind::Periodic,
1,
phi_k[atom_idx].clone(),
jet_k[atom_idx].clone(),
b,
Array2::<f64>::eye(m),
)
.unwrap()
.with_basis_evaluator(Arc::new(PeriodicHarmonicEvaluator::new(m).unwrap()));
atoms.push(atom);
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, k)),
coords_k,
vec![LatentManifold::Circle { period: 1.0 }; k],
AssignmentMode::ordered_beta_bernoulli(0.7, 1.0, false),
)
.unwrap();
let mut term = SaeManifoldTerm::new(atoms, assignment).unwrap();
let mut rho = SaeManifoldRho::new((0.02_f64).ln(), -6.0, vec![Array1::<f64>::zeros(1); k]);
let mut prev_total = f64::INFINITY;
for _ in 0..30 {
let loss = term
.run_joint_fit_arrow_schur(z.view(), &mut rho, None, 1, 1.0, 1.0e-6, 1.0e-6)
.unwrap();
let total = loss.total();
if !total.is_finite() {
break;
}
let denom = prev_total.abs().max(1.0e-12);
let rel = (prev_total - total).abs() / denom;
prev_total = total;
if rel < 1.0e-6 {
break;
}
}
let fitted = term.fitted();
let mut ssr = 0.0;
let mut sst = 0.0;
for i in 0..n {
for j in 0..p {
let r = z[[i, j]] - fitted[[i, j]];
ssr += r * r;
sst += z[[i, j]] * z[[i, j]];
}
}
let r2 = 1.0 - ssr / sst.max(1.0e-12);
assert!(
r2 > 0.5,
"K=2 periodic torus ordered Beta--Bernoulli R² = {r2:.4} (ssr={ssr:.4}, sst={sst:.4}) should be > 0.5 with LSQ-seeded decoder"
);
let assignments = term.assignment.assignments();
let mean_active: f64 = assignments.iter().copied().sum::<f64>() / (n as f64);
assert!(
mean_active > 0.2,
"mean active mass across rows = {mean_active:.4} should exceed 0.2; assignment did not collapse"
);
}
#[test]
pub(crate) fn softmax_k2_periodic_completes_joint_fit_step() {
let n = 64usize;
let p = 4usize;
let k = 2usize;
let m = 3usize;
let mut z = Array2::<f64>::zeros((n, p));
for i in 0..n {
let a = 2.0 * std::f64::consts::PI * (i as f64) / (n as f64);
z[[i, 0]] = a.sin();
z[[i, 1]] = a.cos();
z[[i, 2]] = (2.0 * a).sin();
z[[i, 3]] = (2.0 * a).cos();
}
let evaluator = PeriodicHarmonicEvaluator::new(m).unwrap();
let mut coords_k = vec![Array2::<f64>::zeros((n, 1)); k];
for i in 0..n {
coords_k[0][[i, 0]] = (i as f64) / (n as f64);
coords_k[1][[i, 0]] = ((i as f64) * 2.0 / (n as f64)).rem_euclid(1.0);
}
let mut atoms = Vec::new();
for atom_idx in 0..k {
let (phi, jet) = evaluator.evaluate(coords_k[atom_idx].view()).unwrap();
let b = Array2::<f64>::from_shape_fn((m, p), |(i, j)| {
0.1 * ((i as f64 + 1.0) * (j as f64 + 1.0)).sin()
});
let atom = SaeManifoldAtom::new_with_provided_function_gram(
format!("a_{atom_idx}"),
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
b,
Array2::<f64>::eye(m),
)
.unwrap()
.with_basis_evaluator(Arc::new(PeriodicHarmonicEvaluator::new(m).unwrap()));
atoms.push(atom);
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, k)),
coords_k,
vec![LatentManifold::Circle { period: 1.0 }; k],
AssignmentMode::softmax(0.7),
)
.unwrap();
let mut term = SaeManifoldTerm::new(atoms, assignment).unwrap();
let mut rho = SaeManifoldRho::new(0.0, -6.0, vec![Array1::<f64>::zeros(1); k]);
let loss0 = term
.run_joint_fit_arrow_schur(z.view(), &mut rho, None, 1, 1.0, 1.0e-6, 1.0e-6)
.expect("softmax K=2 must complete first joint-fit step");
assert!(loss0.total().is_finite());
let loss1 = term
.run_joint_fit_arrow_schur(z.view(), &mut rho, None, 1, 1.0, 1.0e-6, 1.0e-6)
.expect("softmax K=2 must complete second joint-fit step");
assert!(loss1.total().is_finite());
}
pub(crate) fn assert_isometry_wiring_matches_fd(
evaluator: Arc<dyn SaeBasisSecondJet>,
coords: Array2<f64>,
) {
let n_obs = coords.nrows();
let latent_dim = coords.ncols();
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let m = phi.ncols();
let p: usize = 3;
let mut decoder = Array2::<f64>::zeros((m, p));
for i in 0..m {
for j in 0..p {
let x = (i as f64) * 0.371 + (j as f64) * 0.193 + 0.5;
decoder[[i, j]] = (x.sin() * 0.9) + 0.1 * ((i + j) as f64).cos();
}
}
let smooth = Array2::<f64>::eye(m);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"iso_wire_test",
SaeAtomBasisKind::Periodic,
latent_dim,
phi.clone(),
jet.clone(),
decoder.clone(),
smooth,
)
.unwrap()
.with_basis_second_jet(evaluator);
let target_slice = PsiSlice::full(n_obs * latent_dim, Some(latent_dim));
let penalty = IsometryPenalty::new_euclidean(target_slice, p);
let rho = Array1::<f64>::zeros(1);
let target_flat: Array1<f64> = coords.iter().copied().collect();
let v0 = penalty.value(target_flat.view(), rho.view());
assert_eq!(v0, IsometryPenalty::DEFAULT_VALUE_ON_MISSING_CACHE);
let g0 = penalty.grad_target(target_flat.view(), rho.view());
assert!(
g0.iter().all(|x| *x == 0.0),
"grad_target without cache must be all zeros, got {g0:?}"
);
let installed_second =
refresh_isometry_caches_from_atom(&penalty, &atom, coords.view()).unwrap();
assert!(
installed_second,
"evaluator must implement second_jet for this oracle to run"
);
let value = penalty.value(target_flat.view(), rho.view());
assert!(
value > 1.0e-6,
"expected non-trivial isometry loss after cache refresh, got {value}"
);
let grad = penalty.grad_target(target_flat.view(), rho.view());
assert_eq!(grad.len(), target_flat.len());
let max_abs = grad.iter().fold(0.0_f64, |acc, x| acc.max(x.abs()));
assert!(
max_abs > 1.0e-6,
"expected non-zero isometry gradient on at least one component, max |grad|={max_abs}"
);
let h_fd = 1.0e-5;
let probe_idx = 0usize; let mut coords_plus = coords.clone();
coords_plus[[0, 0]] += h_fd;
let mut coords_minus = coords.clone();
coords_minus[[0, 0]] -= h_fd;
refresh_isometry_caches_from_atom(&penalty, &atom, coords_plus.view()).unwrap();
let target_plus: Array1<f64> = coords_plus.iter().copied().collect();
let v_plus = penalty.value(target_plus.view(), rho.view());
refresh_isometry_caches_from_atom(&penalty, &atom, coords_minus.view()).unwrap();
let target_minus: Array1<f64> = coords_minus.iter().copied().collect();
let v_minus = penalty.value(target_minus.view(), rho.view());
refresh_isometry_caches_from_atom(&penalty, &atom, coords.view()).unwrap();
let grad_base = penalty.grad_target(target_flat.view(), rho.view());
let fd = (v_plus - v_minus) / (2.0 * h_fd);
let analytic = grad_base[probe_idx];
assert!(
(analytic - fd).abs() <= 1.0e-3 + 1.0e-4 * analytic.abs().max(fd.abs()),
"isometry grad/FD mismatch at coord 0: analytic={analytic:.6e}, fd={fd:.6e}"
);
}
#[test]
pub(crate) fn isometry_wiring_periodic_matches_fd() {
assert_isometry_wiring_matches_fd(
Arc::new(PeriodicHarmonicEvaluator::new(5).unwrap()),
array![[0.12], [0.37], [0.58], [0.81]],
);
}
#[test]
pub(crate) fn isometry_wiring_sphere_matches_fd() {
assert_isometry_wiring_matches_fd(
Arc::new(SphereChartEvaluator),
array![[-0.5, 0.3], [0.2, -1.1], [0.7, 0.9]],
);
}
#[test]
pub(crate) fn isometry_wiring_torus_matches_fd() {
assert_isometry_wiring_matches_fd(
Arc::new(TorusHarmonicEvaluator::new(2, 2).unwrap()),
array![[0.13, 0.42], [0.66, 0.19], [0.88, 0.55]],
);
}
pub(crate) fn warmstart_test_objective() -> SaeManifoldOuterObjective {
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(3).unwrap());
let coords = array![[0.10], [0.35], [0.62], [0.88]];
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.30], [-0.20], [0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(evaluator.clone())
.with_basis_second_jet(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode(
array![[0.9_f64], [0.8], [0.7], [0.6]],
vec![coords],
AssignmentMode::softmax(0.7),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.20_f64], [-0.10], [0.30], [0.05]];
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1)]);
SaeManifoldOuterObjective::new(term, target, None, rho, 8, 1.0, 1.0e-6, 1.0e-6)
}
pub(crate) fn warmstart_test_objective_with_evaluator() -> SaeManifoldOuterObjective {
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(3).unwrap());
let coords = array![[0.10_f64], [0.35], [0.62], [0.88]];
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"periodic",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
array![[0.30_f64], [-0.20], [0.15]],
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_evaluator(evaluator.clone())
.with_basis_second_jet(evaluator);
let assignment = SaeAssignment::from_blocks_with_mode(
array![[0.9_f64], [0.8], [0.7], [0.6]],
vec![coords],
AssignmentMode::softmax(0.7),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[0.20_f64], [-0.10], [0.30], [0.05]];
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1)]);
SaeManifoldOuterObjective::new(term, target, None, rho, 8, 1.0, 1.0e-6, 1.0e-6)
}
pub(crate) fn near_singular_outer_gradient_cache() -> ArrowFactorCache {
ArrowFactorCache {
htt_factors: ArrowFactorSlab::from_blocks(vec![array![[1.0_f64, 0.0], [0.0, 1.0e-7]]]),
htt_factors_undamped: ArrowUndampedFactors::SameAsDamped,
schur_factor: Some(array![[1.0_f64]]),
schur_factor_is_undamped: true,
beta_schur_deflation: None,
joint_hessian_log_det: None,
solver_mode: ArrowSolverMode::Direct,
ridge_t: 0.0,
ridge_beta: 0.0,
htbeta: ArrowHtbetaCache::Disabled { estimated_bytes: 0 },
d: 2,
row_dims: Arc::from(vec![2usize].into_boxed_slice()),
row_offsets: Arc::from(vec![0usize, 2usize].into_boxed_slice()),
k: 1,
manifold_mode_fingerprint: 0,
row_hessian_fingerprint: 0,
pcg_diagnostics: ArrowPcgDiagnostics::default(),
gauge_deflated_directions: 0,
deflated_row_directions: std::sync::Arc::from(Vec::new()),
deflation_row_spectra: std::sync::Arc::from(Vec::new()),
beta_gauge_quotient: None,
}
}
pub(crate) fn diagonal_latent_cache(diagonal: &[f64]) -> ArrowFactorCache {
let dim = diagonal.len();
let mut factor = Array2::<f64>::zeros((dim, dim));
for i in 0..dim {
factor[[i, i]] = diagonal[i].sqrt();
}
ArrowFactorCache {
htt_factors: ArrowFactorSlab::from_blocks(vec![factor]),
htt_factors_undamped: ArrowUndampedFactors::SameAsDamped,
schur_factor: None,
schur_factor_is_undamped: true,
beta_schur_deflation: None,
joint_hessian_log_det: None,
solver_mode: ArrowSolverMode::Direct,
ridge_t: 0.0,
ridge_beta: 0.0,
htbeta: ArrowHtbetaCache::Disabled { estimated_bytes: 0 },
d: dim,
row_dims: Arc::from(vec![dim].into_boxed_slice()),
row_offsets: Arc::from(vec![0usize, dim].into_boxed_slice()),
k: 0,
manifold_mode_fingerprint: 0,
row_hessian_fingerprint: 0,
pcg_diagnostics: ArrowPcgDiagnostics::default(),
gauge_deflated_directions: 0,
deflated_row_directions: std::sync::Arc::from(Vec::new()),
deflation_row_spectra: std::sync::Arc::from(Vec::new()),
beta_gauge_quotient: None,
}
}
#[test]
pub(crate) fn outer_gradient_solver_rejects_near_singular_cache_without_matching_gauge() {
let cache = near_singular_outer_gradient_cache();
let obj = warmstart_test_objective();
let conditioning_err = match SaeManifoldTerm::outer_gradient_conditioning_error(&cache) {
Err(err) => err.to_string(),
Ok(()) => panic!("near-singular cache must trip the pivot-ratio conditioning gate"),
};
assert!(
conditioning_err.contains("joint Hessian numerically singular"),
"conditioning gate must name the ill-conditioned joint Hessian; got: {conditioning_err}"
);
assert!(
conditioning_err.contains("min/max pivot ratio") && conditioning_err.contains("floor"),
"conditioning gate must report the pivot ratio and floor; got: {conditioning_err}"
);
let err = match obj
.term
.outer_gradient_arrow_solver(&cache, &obj.current_rho.lambda_smooth_vec().unwrap())
{
Err(err) => err,
Ok(..) => panic!("near-singular criterion factor without a matching gauge must reject"),
};
assert!(
matches!(err, OuterGradientError::NonIdentifiable { .. }),
"no-deflatable-direction rejection must be the NonIdentifiable diagnosis; got: {err}"
);
let err = err.to_string();
assert!(
err.contains("no deflatable gauge/decoder-null direction"),
"guard error must name the absent deflation candidate; got: {err}"
);
}
pub(crate) fn rank_deficient_euclidean_outer_gradient_objective() -> SaeManifoldOuterObjective {
let coords = array![[-0.7_f64], [-0.2], [0.3], [0.8]];
let n = coords.nrows();
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 decoder = array![[1.0_f64, 0.0], [0.5, 0.0]];
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"euclidean_line",
SaeAtomBasisKind::EuclideanPatch,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(2),
)
.unwrap();
let assignment = SaeAssignment::from_blocks_with_mode(
array![[0.9_f64], [0.8], [0.7], [0.6]],
vec![coords],
AssignmentMode::softmax(0.7),
)
.unwrap();
let term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[-1.0_f64, -2.0], [-0.3, -0.6], [0.4, 0.8], [1.1, 2.2]];
let rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1)]);
SaeManifoldOuterObjective::new(term, target, None, rho, 8, 1.0, 1.0e-6, 1.0e-6)
}
pub(crate) fn rank_deficient_beta_outer_gradient_cache() -> ArrowFactorCache {
let htt = ArrowFactorSlab::from_blocks(vec![
array![[1.0_f64]],
array![[1.0_f64]],
array![[1.0_f64]],
array![[1.0_f64]],
]);
let schur = array![
[1.0_f64, 0.0, 0.0, 0.0],
[0.0, 1.0e-7, 0.0, 0.0],
[0.0, 0.0, 1.0, 0.0],
[0.0, 0.0, 0.0, 1.0e-7],
];
ArrowFactorCache {
htt_factors: htt,
htt_factors_undamped: ArrowUndampedFactors::SameAsDamped,
schur_factor: Some(schur),
schur_factor_is_undamped: true,
beta_schur_deflation: None,
joint_hessian_log_det: None,
solver_mode: ArrowSolverMode::Direct,
ridge_t: 0.0,
ridge_beta: 0.0,
htbeta: ArrowHtbetaCache::Dense {
blocks: Arc::from(
vec![
Array2::<f64>::zeros((1, 4)),
Array2::<f64>::zeros((1, 4)),
Array2::<f64>::zeros((1, 4)),
Array2::<f64>::zeros((1, 4)),
]
.into_boxed_slice(),
),
estimated_bytes: 0,
},
d: 4,
row_dims: Arc::from(vec![1usize, 1usize, 1usize, 1usize].into_boxed_slice()),
row_offsets: Arc::from(vec![0usize, 1usize, 2usize, 3usize, 4usize].into_boxed_slice()),
k: 4,
manifold_mode_fingerprint: 0,
row_hessian_fingerprint: 0,
pcg_diagnostics: ArrowPcgDiagnostics::default(),
gauge_deflated_directions: 0,
deflated_row_directions: std::sync::Arc::from(Vec::new()),
deflation_row_spectra: std::sync::Arc::from(Vec::new()),
beta_gauge_quotient: None,
}
}
#[test]
pub(crate) fn outer_gradient_solver_deflates_rank_deficient_decoder_beta_null() {
let obj = rank_deficient_euclidean_outer_gradient_objective();
let cache = rank_deficient_beta_outer_gradient_cache();
assert!(
SaeManifoldTerm::outer_gradient_conditioning_error(&cache).is_err(),
"fixture must be sub-floor singular so the conditioning path engages"
);
let solver = obj
.term
.outer_gradient_arrow_solver(&cache, &obj.current_rho.lambda_smooth_vec().unwrap())
.expect("rank-deficient decoder β-null must be deflated, not rejected (#1051/#1273)");
let beta_null_rhs = array![0.0_f64, 0.0, 0.0, 1.0]; let rhs_t = Array1::<f64>::zeros(cache.delta_t_len());
let plain = cache
.full_inverse_apply(rhs_t.view(), beta_null_rhs.view())
.expect("plain solve")
.1;
let deflated = solver
.solve(rhs_t.view(), beta_null_rhs.view())
.expect("deflated solve")
.beta;
assert!(
plain[3].abs() > 1.0e13,
"plain near-null β solve must explode; got {}",
plain[3]
);
assert!(
deflated.iter().all(|v| v.is_finite()) && deflated[3].abs() < 10.0,
"deflated near-null β solve must be bounded at the Hessian scale; got {deflated:?}"
);
}
#[test]
pub(crate) fn outer_gradient_internal_invariant_is_typed_1436() {
let ill_conditioned = OuterGradientError::IllConditioned {
reason: "near-singular joint Hessian".to_string(),
};
let non_identifiable = OuterGradientError::NonIdentifiable {
reason: "gauge-degenerate direction".to_string(),
};
let internal = OuterGradientError::InternalInvariant {
reason: "shape mismatch".to_string(),
};
assert!(ill_conditioned.to_string().contains("ill-conditioned"));
assert!(non_identifiable.to_string().contains("non-identifiable"));
assert!(
internal.to_string().contains("internal invariant"),
"InternalInvariant Display must name the class; got: {}",
internal
);
}
#[test]
pub(crate) fn seed_inner_state_accepts_empty_beta_as_noslot() {
let mut obj = warmstart_test_objective();
let empty: Array1<f64> = Array1::zeros(0);
let outcome = obj
.seed_inner_state(&empty)
.expect("empty-β seed must be accepted as a no-op, not rejected (gam#577/#579)");
assert!(
matches!(outcome, SeedOutcome::NoSlot),
"empty-β seed must report NoSlot (proceed cold); got {outcome:?}"
);
}
#[test]
pub(crate) fn seed_inner_state_installs_and_reuses_matching_beta() {
let mut source = warmstart_test_objective();
let source_rho = source.baseline_rho.clone();
source
.term
.penalized_quasi_laplace_criterion_with_cache(
source.target.view(),
&source_rho,
source.registry.as_ref(),
source.inner_max_iter,
source.learning_rate,
source.ridge_ext_coord,
source.ridge_beta,
)
.expect("source continuation state must have finite converged evidence");
let seed = source.term.flatten_beta();
let mut obj = warmstart_test_objective();
let dim = obj.term.beta_dim();
let pristine = obj.term.flatten_beta();
assert_eq!(seed.len(), dim, "source β must match the target layout");
assert!(
(&seed - &pristine).iter().any(|d| d.abs() > 1e-6),
"converged continuation β must differ from the pristine target β for the reuse check"
);
let outcome = obj
.seed_inner_state(&seed)
.expect("a length-matching β must install");
assert!(
matches!(outcome, SeedOutcome::Installed),
"matching β must report Installed; got {outcome:?}"
);
obj.inner_max_iter = 0;
let rho_flat = obj.baseline_rho.to_flat();
let eval =
OuterObjective::eval(&mut obj, &rho_flat).expect("eval at the warm-started β must succeed");
let hint = eval
.inner_beta_hint
.expect("the SAE objective must publish inner_beta_hint for continuation reuse");
assert_eq!(
hint.len(),
dim,
"published hint must have decoder dimension"
);
for (i, (&h, &s)) in hint.iter().zip(seed.iter()).enumerate() {
assert!(
(h - s).abs() < 1e-12,
"warm-started β must be reused verbatim by the inner solve at coord {i}: \
hint {h} != seed {s} (gam#577/#579)"
);
}
}
#[test]
pub(crate) fn seed_inner_state_rejects_wrong_length_populated_beta() {
let mut obj = warmstart_test_objective();
let dim = obj.term.beta_dim();
let wrong: Array1<f64> = Array1::zeros(dim + 1);
let err = obj
.seed_inner_state(&wrong)
.expect_err("a populated β of the wrong length must be rejected");
match err {
EstimationError::RemlOptimizationFailed(msg) => {
assert!(
msg.contains("decoder dim"),
"error must name the decoder-dim mismatch; got: {msg}"
);
}
other => panic!("expected RemlOptimizationFailed, got {other:?}"),
}
}
#[test]
pub(crate) fn reference_function_gram_is_fixed_and_has_exact_trace_form() {
let n = 4usize;
let m = 3usize;
let p = 2usize;
let phi = Array2::<f64>::zeros((n, m));
let jet = Array3::<f64>::zeros((n, m, 1));
let decoder = Array2::from_shape_vec((m, p), vec![0.3, -0.2, 1.1, 0.4, -0.7, 0.9]).unwrap();
let gram = gam_terms::basis::create_difference_penalty_matrix(m, 2, None).unwrap();
let mut atom = SaeManifoldAtom::new_with_provided_function_gram(
"fixed-reference",
SaeAtomBasisKind::EuclideanPatch,
1,
phi,
jet,
decoder.clone(),
gram.clone(),
)
.unwrap();
let frozen = atom.smooth_penalty().clone();
atom.decoder_coefficients.mapv_inplace(|value| value * 7.0);
atom.basis_jacobian.fill(13.0);
assert_eq!(atom.smooth_penalty(), &frozen);
assert_eq!(
atom.reference_roughness_kind(),
SaeReferenceRoughnessKind::ProvidedFunctionGram
);
let trace_form: f64 = (0..p)
.map(|output| {
let column = decoder.column(output);
column.dot(&gram.dot(&column))
})
.sum();
let bt_s_b = decoder.t().dot(&gram).dot(&decoder);
let matrix_form = bt_s_b.diag().sum();
assert_abs_diff_eq!(trace_form, matrix_form, epsilon = 1.0e-12);
}
pub(crate) fn gamma_fd_tiny_fixture() -> (SaeManifoldTerm, Array2<f64>, SaeManifoldRho) {
let n = 10usize;
let p = 3usize;
let k_atoms = 2usize;
let m = 3usize;
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(m).unwrap());
let mut logits = Array2::<f64>::zeros((n, k_atoms));
let mut coords = vec![Array2::<f64>::zeros((n, 1)), Array2::<f64>::zeros((n, 1))];
let weights = [
[
[0.10, -0.05, 0.03],
[0.35, -0.20, 0.12],
[-0.16, 0.18, 0.08],
],
[
[-0.08, 0.04, 0.06],
[0.22, 0.10, -0.18],
[0.11, -0.24, 0.15],
],
];
let mut target = Array2::<f64>::zeros((n, p));
for row in 0..n {
let phase = (row as f64 + 0.35) / n as f64;
coords[0][[row, 0]] = phase;
coords[1][[row, 0]] = (phase + 0.21).fract();
logits[[row, 0]] = if row % 2 == 0 { 0.8 } else { -0.6 };
let assignments = softmax_row(logits.row(row), 0.9);
for atom in 0..k_atoms {
let theta = std::f64::consts::TAU * coords[atom][[row, 0]];
let basis = [1.0, theta.sin(), theta.cos()];
for out_col in 0..p {
for basis_col in 0..m {
target[[row, out_col]] +=
assignments[atom] * basis[basis_col] * weights[atom][basis_col][out_col];
}
}
}
}
let mut atoms = Vec::with_capacity(k_atoms);
for atom in 0..k_atoms {
let (phi, jet) = evaluator.evaluate(coords[atom].view()).unwrap();
let decoder = Array2::from_shape_fn((m, p), |(basis_col, out_col)| {
weights[atom][basis_col][out_col]
});
atoms.push(
SaeManifoldAtom::new_with_provided_function_gram(
format!("gamma_{atom}"),
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(m),
)
.unwrap()
.with_basis_second_jet(evaluator.clone()),
);
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coords,
vec![LatentManifold::Circle { period: 1.0 }; k_atoms],
AssignmentMode::softmax(0.9),
)
.unwrap();
let term = SaeManifoldTerm::new(atoms, assignment).unwrap();
let rho = SaeManifoldRho::new(
-6.0,
-6.0,
vec![Array1::from_vec(vec![-6.0]), Array1::from_vec(vec![-6.0])],
);
(term, target, rho)
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct FiniteDifferenceStratumCertificate {
row_dims: Vec<usize>,
row_offsets: Vec<usize>,
beta_dim: usize,
manifold_mode_fingerprint: u64,
solver_mode: ArrowSolverMode,
gauge_deflated_directions: usize,
deflated_per_row: Vec<usize>,
row_spectral_conditioning: Vec<Option<Vec<RowSpectralConditioning>>>,
beta_schur_deflated: Option<Vec<bool>>,
beta_gauge_rank: usize,
eigen_derivative_route: EigenDerivativeRoute,
min_row_pivot_branch: PivotBranch,
min_schur_pivot_branch: PivotBranch,
min_pivot_branch: PivotBranch,
max_pivot_branch: PivotBranch,
}
impl FiniteDifferenceStratumCertificate {
pub(crate) fn from_arrow_cache(cache: &ArrowFactorCache) -> Self {
let min_pivot = arrow_factor_min_pivot(cache);
Self {
row_dims: cache.row_dims.to_vec(),
row_offsets: cache.row_offsets.to_vec(),
beta_dim: cache.k,
manifold_mode_fingerprint: cache.manifold_mode_fingerprint,
solver_mode: cache.solver_mode,
gauge_deflated_directions: cache.gauge_deflated_directions,
deflated_per_row: cache.deflated_row_directions.iter().map(Vec::len).collect(),
row_spectral_conditioning: cache
.deflation_row_spectra
.iter()
.map(|spectrum| {
spectrum
.as_ref()
.map(|spectrum| spectrum.conditioning.to_vec())
})
.collect(),
beta_schur_deflated: cache
.beta_schur_deflation
.as_ref()
.map(|spectrum| spectrum.deflated.to_vec()),
beta_gauge_rank: cache
.beta_gauge_quotient
.as_ref()
.map_or(0, |quotient| quotient.directions.len()),
eigen_derivative_route: BranchCertificate::from_arrow_cache(
cache,
MajorizerAnchorMode::FrozenAnchor,
)
.eigen_derivative_route(),
min_row_pivot_branch: classify_fd_pivot(min_pivot.min_row_pivot),
min_schur_pivot_branch: classify_fd_pivot(min_pivot.min_schur_pivot),
min_pivot_branch: classify_fd_pivot(min_pivot.min_pivot),
max_pivot_branch: classify_fd_pivot(arrow_factor_max_pivot(cache)),
}
}
pub(crate) fn changed_fields(&self, endpoint: &Self) -> Vec<&'static str> {
let mut changed = Vec::new();
if self.row_dims != endpoint.row_dims {
changed.push("row_dims");
}
if self.row_offsets != endpoint.row_offsets {
changed.push("row_offsets");
}
if self.beta_dim != endpoint.beta_dim {
changed.push("beta_dim");
}
if self.manifold_mode_fingerprint != endpoint.manifold_mode_fingerprint {
changed.push("manifold_mode");
}
if self.solver_mode != endpoint.solver_mode {
changed.push("solver_mode");
}
if self.gauge_deflated_directions != endpoint.gauge_deflated_directions {
changed.push("gauge_deflated_directions");
}
if self.deflated_per_row != endpoint.deflated_per_row {
changed.push("deflated_per_row");
}
if self.row_spectral_conditioning != endpoint.row_spectral_conditioning {
changed.push("row_spectral_conditioning");
}
if self.beta_schur_deflated != endpoint.beta_schur_deflated {
changed.push("beta_schur_deflated");
}
if self.beta_gauge_rank != endpoint.beta_gauge_rank {
changed.push("beta_gauge_rank");
}
if self.eigen_derivative_route != endpoint.eigen_derivative_route {
changed.push("eigen_derivative_route");
}
if self.min_row_pivot_branch != endpoint.min_row_pivot_branch {
changed.push("min_row_pivot_branch");
}
if self.min_schur_pivot_branch != endpoint.min_schur_pivot_branch {
changed.push("min_schur_pivot_branch");
}
if self.min_pivot_branch != endpoint.min_pivot_branch {
changed.push("min_pivot_branch");
}
if self.max_pivot_branch != endpoint.max_pivot_branch {
changed.push("max_pivot_branch");
}
changed
}
pub(crate) fn assert_same_stratum(&self, label: &str, endpoint: &Self) {
assert_eq!(
self.eigen_derivative_route,
EigenDerivativeRoute::IndividualEigenpairs,
"{label}: finite-difference center has an unresolved spectral invariant-subspace block"
);
assert_eq!(
endpoint.eigen_derivative_route,
EigenDerivativeRoute::IndividualEigenpairs,
"{label}: finite-difference endpoint has an unresolved spectral invariant-subspace block"
);
let changed = self.changed_fields(endpoint);
assert!(
changed.is_empty(),
"{label}: finite-difference endpoint crossed a nondifferentiable structural stratum; \
changed_fields={changed:?}\ncenter={self:#?}\nendpoint={endpoint:#?}"
);
}
}
fn classify_fd_pivot(pivot: Option<f64>) -> PivotBranch {
match pivot {
None => PivotBranch::Missing,
Some(value) if !value.is_finite() => PivotBranch::NonFinite,
Some(value) if value > 0.0 => PivotBranch::Positive,
Some(_) => PivotBranch::NonPositive,
}
}
#[derive(Clone, Debug)]
pub(crate) struct FixedStateLogdetSample {
pub(crate) value: f64,
pub(crate) stratum: FiniteDifferenceStratumCertificate,
}
pub(crate) fn fixed_state_logdet_sample(
mut term: SaeManifoldTerm,
target: &Array2<f64>,
rho: &SaeManifoldRho,
) -> FixedStateLogdetSample {
let (_value, _loss, cache) = term
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
rho,
None,
0,
0.4,
1.0e-6,
1.0e-6,
)
.expect("fixed-state cache");
let value = arrow_log_det_from_cache(&cache).expect("fixed-state authoritative joint logdet");
let stratum = FiniteDifferenceStratumCertificate::from_arrow_cache(&cache);
FixedStateLogdetSample { value, stratum }
}
pub(crate) fn certified_central_logdet_difference(
label: &str,
center: &FiniteDifferenceStratumCertificate,
plus: FixedStateLogdetSample,
minus: FixedStateLogdetSample,
step: f64,
) -> f64 {
assert!(
step.is_finite() && step > 0.0,
"{label}: central-difference step must be finite and positive, got {step}"
);
center.assert_same_stratum(&format!("{label} (+h)"), &plus.stratum);
center.assert_same_stratum(&format!("{label} (-h)"), &minus.stratum);
(plus.value - minus.value) / (2.0 * step)
}
fn floor_crossing_logdet_sample_2398(ratio: f64) -> FixedStateLogdetSample {
assert!(ratio.is_finite() && ratio > 0.0);
let relative_floor = gam_solve::arrow_schur::SPECTRAL_DEFLATION_REL_FLOOR;
let spectral_floor = relative_floor * 4.0;
let mut system = ArrowSchurSystem::new(1, 3, 0);
system.rows[0].htt = array![
[4.0_f64, 0.0, 0.0],
[0.0, ratio * spectral_floor, 0.0],
[0.0, 0.0, -1.0],
];
SaeManifoldTerm::ensure_row_gauge_deflation_for_quasi_laplace(&mut system);
let options = ArrowSolveOptions::direct()
.with_gpu_policy(gam_gpu::GpuPolicy::Off)
.with_evidence_unit_deflation(relative_floor);
let (_, _, cache) = solve_arrow_newton_step_with_options(&system, 2.0, 0.0, &options)
.expect("floor-crossing fixture must produce an evidence cache");
let value =
arrow_log_det_from_cache(&cache).expect("floor-crossing fixture must have a logdet");
let stratum = FiniteDifferenceStratumCertificate::from_arrow_cache(&cache);
FixedStateLogdetSample { value, stratum }
}
#[test]
fn certified_central_logdet_difference_refuses_floor_clamp_crossing_2398() {
let minus_ratio = 0.995_f64;
let plus_ratio = 1.05_f64;
let center_ratio = (minus_ratio + plus_ratio) / 2.0;
let step = (plus_ratio - minus_ratio) / 2.0;
let center = floor_crossing_logdet_sample_2398(center_ratio);
let plus = floor_crossing_logdet_sample_2398(plus_ratio);
let minus = floor_crossing_logdet_sample_2398(minus_ratio);
let expected_raw = vec![
RowSpectralConditioning::UnitDeflated,
RowSpectralConditioning::Raw,
RowSpectralConditioning::Raw,
];
let expected_clamped = vec![
RowSpectralConditioning::UnitDeflated,
RowSpectralConditioning::FloorClamped,
RowSpectralConditioning::Raw,
];
assert_eq!(
center.stratum.row_spectral_conditioning,
vec![Some(expected_raw.clone())],
);
assert_eq!(
plus.stratum.row_spectral_conditioning,
vec![Some(expected_raw)],
);
assert_eq!(
minus.stratum.row_spectral_conditioning,
vec![Some(expected_clamped)],
);
assert!(center.stratum.changed_fields(&plus.stratum).is_empty());
assert_eq!(
center.stratum.changed_fields(&minus.stratum),
vec!["row_spectral_conditioning"],
);
let panic = std::panic::catch_unwind(move || {
certified_central_logdet_difference(
"spectral-floor stratum #2398",
¢er.stratum,
plus,
minus,
step,
)
})
.expect_err("a floor-clamp crossing must refuse before producing a quotient");
let message = panic
.downcast_ref::<String>()
.map(String::as_str)
.or_else(|| panic.downcast_ref::<&str>().copied())
.expect("refusal panic must carry a string message");
assert!(
message.contains("finite-difference endpoint crossed a nondifferentiable structural stratum")
);
assert!(message.contains("changed_fields=[\"row_spectral_conditioning\"]"));
assert!(message.contains("spectral-floor stratum #2398 (-h)"));
}
#[derive(Clone, Copy, Debug)]
pub(crate) enum FdBranchRegime {
Smooth { ratio: f64 },
RoundoffDominated { coarse_gap: f64, floor: f64 },
}
pub(crate) fn certified_branch_stable_central_difference(
label: &str,
center: &FiniteDifferenceStratumCertificate,
step: f64,
sample: impl Fn(f64) -> FixedStateLogdetSample,
) -> (f64, FdBranchRegime) {
assert!(
step.is_finite() && step > 0.0,
"{label}: central-difference step must be finite and positive, got {step}"
);
let mut magnitude = 0.0_f64;
let mut quotients = [0.0_f64; 3];
for (index, scale) in [1.0_f64, 0.5, 0.25].into_iter().enumerate() {
let h = step * scale;
let plus = sample(h);
let minus = sample(-h);
center.assert_same_stratum(&format!("{label} (+{scale}h)"), &plus.stratum);
center.assert_same_stratum(&format!("{label} (-{scale}h)"), &minus.stratum);
magnitude = magnitude.max(plus.value.abs()).max(minus.value.abs());
quotients[index] = (plus.value - minus.value) / (2.0 * h);
}
let coarse_gap = (quotients[0] - quotients[1]).abs();
let fine_gap = (quotients[1] - quotients[2]).abs();
let floor = f64::EPSILON * magnitude / (0.25 * step);
if coarse_gap <= 100.0 * floor {
return (
quotients[0],
FdBranchRegime::RoundoffDominated { coarse_gap, floor },
);
}
let ratio = fine_gap / coarse_gap;
assert!(
ratio <= 0.35,
"{label}: central-difference gaps must fall as h² across a branch-stable \
stencil; coarse(h={step:.3e}) {coarse_gap:.6e}, fine(h/2) {fine_gap:.6e}, \
ratio {ratio:.4} (predicted 0.25; O(h) kink gives 0.5, O(1) jump gives 1). \
The stratum certificate passed, so this is nonsmoothness the classifier \
does not label."
);
(quotients[0], FdBranchRegime::Smooth { ratio })
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum FdAnchorDeflation {
NoRowDeflates,
SomeRowDeflates,
Unconstrained,
}
#[derive(Clone, Copy, Debug)]
pub(crate) struct FdAnchorRegime {
pub(crate) deflation: FdAnchorDeflation,
pub(crate) reportable_derivative_branch: bool,
}
impl FdAnchorRegime {
pub(crate) fn smooth_majorizer_branch() -> Self {
Self {
deflation: FdAnchorDeflation::Unconstrained,
reportable_derivative_branch: true,
}
}
pub(crate) fn undeflated() -> Self {
Self {
deflation: FdAnchorDeflation::NoRowDeflates,
reportable_derivative_branch: false,
}
}
pub(crate) fn any_maximum() -> Self {
Self {
deflation: FdAnchorDeflation::Unconstrained,
reportable_derivative_branch: false,
}
}
pub(crate) fn deflated() -> Self {
Self {
deflation: FdAnchorDeflation::SomeRowDeflates,
reportable_derivative_branch: false,
}
}
}
pub(crate) struct CertifiedFdAnchor {
pub(crate) term: SaeManifoldTerm,
pub(crate) rho: SaeManifoldRho,
pub(crate) cache: ArrowFactorCache,
pub(crate) stratum: FiniteDifferenceStratumCertificate,
}
pub(crate) struct FdAnchorCandidate {
pub(crate) description: String,
pub(crate) rho: SaeManifoldRho,
pub(crate) term: SaeManifoldTerm,
pub(crate) converge_iters: usize,
pub(crate) converge_tolerance: f64,
pub(crate) decisive_mix: Option<(Array2<f64>, f64)>,
}
fn classify_fd_anchor_candidate(
target: &Array2<f64>,
regime: FdAnchorRegime,
candidate: FdAnchorCandidate,
) -> Result<CertifiedFdAnchor, String> {
let FdAnchorCandidate {
rho,
mut term,
converge_iters,
converge_tolerance,
decisive_mix,
..
} = candidate;
if converge_iters > 0 {
term.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
converge_iters,
0.4,
converge_tolerance,
converge_tolerance,
)
.map_err(|error| format!("inner solve refused to converge: {error}"))?;
}
if let Some((decisive, weight)) = decisive_mix {
if decisive.dim() != term.assignment.logits.dim() {
return Err(format!(
"decisive logit pattern {:?} does not match the assignment layout {:?}",
decisive.dim(),
term.assignment.logits.dim()
));
}
term.assignment.logits = &term.assignment.logits * (1.0 - weight) + &decisive * weight;
}
let (value, loss, cache) = term
.penalized_quasi_laplace_criterion_with_cache(
target.view(),
&rho,
None,
0,
0.4,
1.0e-6,
1.0e-6,
)
.map_err(|error| format!("criterion refused the frozen state: {error}"))?;
if !(value.is_finite() && loss.total().is_finite()) {
return Err(format!(
"frozen state priced non-finitely (value={value}, loss={})",
loss.total()
));
}
let deflated_rows = cache
.deflated_row_directions
.iter()
.filter(|directions| !directions.is_empty())
.count();
match regime.deflation {
FdAnchorDeflation::NoRowDeflates if deflated_rows > 0 => {
return Err(format!(
"{deflated_rows} row(s) deflate; regime requires none"
));
}
FdAnchorDeflation::SomeRowDeflates if deflated_rows == 0 => {
return Err("no row deflates; regime requires at least one".to_string());
}
FdAnchorDeflation::NoRowDeflates
| FdAnchorDeflation::SomeRowDeflates
| FdAnchorDeflation::Unconstrained => {}
}
if regime.reportable_derivative_branch {
let certificate =
BranchCertificate::from_arrow_cache(&cache, MajorizerAnchorMode::FrozenAnchor);
certificate
.assert_derivative_reportable()
.map_err(|error| format!("derivative branch is not reportable: {error}"))?;
if let Some(record) = certificate
.kink_branches
.iter()
.find(|record| record.branch == DualKinkBranch::Tie)
{
return Err(format!("majorizer kink branch is tied: {record:?}"));
}
for (name, branch) in [
("min row pivot", certificate.min_row_pivot_branch),
("min pivot", certificate.min_pivot_branch),
("max pivot", certificate.max_pivot_branch),
] {
if branch != PivotBranch::Positive {
return Err(format!("{name} branch is {branch:?}, not Positive"));
}
}
if certificate.beta_dim > 0 && certificate.min_schur_pivot_branch != PivotBranch::Positive {
return Err(format!(
"min Schur pivot branch is {:?}, not Positive",
certificate.min_schur_pivot_branch
));
}
}
let stratum = FiniteDifferenceStratumCertificate::from_arrow_cache(&cache);
Ok(CertifiedFdAnchor {
term,
rho,
cache,
stratum,
})
}
pub(crate) fn certified_fd_anchor(
label: &str,
target: &Array2<f64>,
regime: FdAnchorRegime,
candidates: Vec<FdAnchorCandidate>,
) -> CertifiedFdAnchor {
assert!(
!candidates.is_empty(),
"{label}: an anchor family must declare at least one candidate"
);
let mut rejections = Vec::with_capacity(candidates.len());
for candidate in candidates {
let description = candidate.description.clone();
match classify_fd_anchor_candidate(target, regime, candidate) {
Ok(anchor) => {
eprintln!("{label}: anchor certified at {description} ({regime:?})");
return anchor;
}
Err(reason) => rejections.push(format!(" {description}: {reason}")),
}
}
panic!(
"{label}: no member of the declared anchor family certifies {regime:?}. \
The regime the asserted derivative needs does not exist anywhere on this \
family, so widening the family is a fixture decision and weakening the \
regime would change what is proved. Rejections:\n{}",
rejections.join("\n")
);
}
pub(crate) fn decisive_logit_homotopy(
converged: &SaeManifoldTerm,
rho: &SaeManifoldRho,
decisive: &Array2<f64>,
) -> Vec<FdAnchorCandidate> {
let base = converged.assignment.logits.clone();
assert_eq!(
base.dim(),
decisive.dim(),
"decisive logit pattern must match the fixture's assignment layout"
);
DECISIVE_HOMOTOPY_WEIGHTS
.into_iter()
.map(|s| {
let mut term = converged.clone();
term.assignment.logits = &base * (1.0 - s) + decisive * s;
FdAnchorCandidate {
description: format!("decisive-logit homotopy s={s:.2}"),
rho: rho.clone(),
term,
converge_iters: 0,
converge_tolerance: DEFAULT_ANCHOR_CONVERGE_TOLERANCE,
decisive_mix: None,
}
})
.collect()
}
const DECISIVE_HOMOTOPY_WEIGHTS: [f64; 7] = [1.0, 0.85, 0.7, 0.55, 0.4, 0.25, 0.1];
pub(crate) fn rho_ladder_family(
term: &SaeManifoldTerm,
rhos: Vec<(String, SaeManifoldRho)>,
converge_iters: usize,
) -> Vec<FdAnchorCandidate> {
rho_ladder_family_with_tolerance(
term,
rhos,
converge_iters,
DEFAULT_ANCHOR_CONVERGE_TOLERANCE,
)
}
pub(crate) fn rho_ladder_family_with_tolerance(
term: &SaeManifoldTerm,
rhos: Vec<(String, SaeManifoldRho)>,
converge_iters: usize,
converge_tolerance: f64,
) -> Vec<FdAnchorCandidate> {
rhos.into_iter()
.map(|(description, rho)| FdAnchorCandidate {
description,
rho,
term: term.clone(),
converge_iters,
converge_tolerance,
decisive_mix: None,
})
.collect()
}
const DEFAULT_ANCHOR_CONVERGE_TOLERANCE: f64 = 1.0e-6;
pub(crate) fn sparse_lift_ladder(
base: &SaeManifoldRho,
lifts: &[f64],
) -> Vec<(String, SaeManifoldRho)> {
lifts
.iter()
.map(|&lift| {
let mut rho = base.clone();
rho.log_lambda_sparse = lift;
(format!("log_lambda_sparse={lift:.2}"), rho)
})
.collect()
}
pub(crate) fn smoothing_and_decisive_family(
fixture: impl Fn() -> (SaeManifoldTerm, SaeManifoldRho),
smooth_levels: &[(f64, f64)],
converge_iters: usize,
) -> Vec<FdAnchorCandidate> {
let mut candidates = Vec::with_capacity(smooth_levels.len() * DECISIVE_HOMOTOPY_WEIGHTS.len());
for &(first, second) in smooth_levels {
for s in DECISIVE_HOMOTOPY_WEIGHTS {
let (term, mut rho) = fixture();
assert_eq!(
rho.log_lambda_smooth.len(),
2,
"the declared smoothing ladder is written for two smooth blocks"
);
rho.log_lambda_smooth = vec![first, second];
let decisive = decisive_logit_pattern(&term);
candidates.push(FdAnchorCandidate {
description: format!(
"log_lambda_smooth=[{first:.2}, {second:.2}], decisive-logit homotopy s={s:.2}"
),
rho,
term,
converge_iters,
converge_tolerance: DEFAULT_ANCHOR_CONVERGE_TOLERANCE,
decisive_mix: Some((decisive, s)),
});
}
}
candidates
}
pub(crate) fn decisive_logit_pattern(term: &SaeManifoldTerm) -> Array2<f64> {
let mut logits = term.assignment.logits.clone();
for row in 0..term.n_obs() {
let center = 0.05 * (row as f64);
let margin = 1.55 + 0.04 * (row as f64);
let (first, second) = if row % 2 == 0 {
(center + margin, center - margin)
} else {
(center - 0.85 * margin, center + 0.85 * margin)
};
logits[[row, 0]] = first;
if logits.ncols() > 1 {
logits[[row, 1]] = second;
}
}
logits
}