#![cfg(test)]
use gam_linalg::faer_ndarray::fast_ata;
use super::*;
use ndarray::array;
pub(crate) fn real_data_torus_seed_term(
z: ArrayView2<'_, f64>,
k: usize,
num_harmonics: usize,
) -> SaeManifoldTerm {
let n = z.nrows();
let evaluator = Arc::new(TorusHarmonicEvaluator::new(2, num_harmonics).unwrap());
let basis_kinds = vec![SaeAtomBasisKind::Periodic; k];
let atom_dims = vec![2usize; k];
let seed_coords = sae_pca_seed_initial_coords(z, &basis_kinds, &atom_dims).unwrap();
let mut atoms = Vec::with_capacity(k);
let mut coords_blocks = Vec::with_capacity(k);
let mut manifolds = Vec::with_capacity(k);
for atom_idx in 0..k {
let coords = seed_coords.slice(s![atom_idx, .., 0..2]).to_owned();
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let m = phi.ncols();
let mut xtx = fast_ata(&phi);
for i in 0..m {
xtx[[i, i]] += 1.0e-8;
}
let xtz = fast_atb(&phi, &z.to_owned());
let decoder = xtx.cholesky(Side::Lower).unwrap().solve_mat(&xtz);
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"torus",
SaeAtomBasisKind::Periodic,
2,
phi,
jet,
decoder,
Array2::<f64>::eye(m),
)
.unwrap()
.with_basis_evaluator(evaluator.clone());
atoms.push(atom);
coords_blocks.push(coords);
manifolds.push(LatentManifold::Product(vec![
LatentManifold::Circle { period: 1.0 },
LatentManifold::Circle { period: 1.0 },
]));
}
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::from_elem((n, k), 0.0),
coords_blocks,
manifolds,
AssignmentMode::softmax(1.0),
)
.unwrap();
SaeManifoldTerm::new(atoms, assignment).unwrap()
}
#[test]
pub(crate) fn olmo_real_curvature_anchor_is_positive_definite() {
let path = olmo_fixture_path("olmo_mixedlayer_pca64_768.npy");
let z = read_npy_f32_2d(&path);
assert_eq!(z.dim(), (768, 64), "real OLMo fixture shape");
let z_train = z.slice(s![..160, ..]).to_owned();
let k = 2usize;
let mut term = real_data_torus_seed_term(z_train.view(), k, 3);
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![0.0, 0.0]; k]);
let registry = SaeManifoldOuterObjective::new(
term.clone(),
z_train.clone(),
None,
rho.clone(),
0,
0.04,
1.0e-6,
1.0e-6,
)
.registry;
use gam_linalg::faer_ndarray::FaerEigh;
let sys = term
.assemble_arrow_schur(z_train.view(), &rho, registry.as_ref())
.expect("assemble raw curvature anchor");
let mut min_raw_eig = f64::INFINITY;
let mut max_raw_eig = 0.0_f64;
let mut indefinite_rows = 0usize;
let mut total_neg_dirs = 0usize;
for block in &sys.rows {
let d = block.htt.nrows();
if d == 0 {
continue;
}
let mut sym = Array2::<f64>::zeros((d, d));
for i in 0..d {
for j in 0..d {
sym[[i, j]] = 0.5 * (block.htt[[i, j]] + block.htt[[j, i]]);
}
}
let (evals, _) = sym.eigh(faer::Side::Lower).unwrap();
let max_abs = evals.iter().fold(0.0_f64, |a, &v| a.max(v.abs())).max(1.0);
let neg_floor = -1.0e-8 * max_abs;
let row_min = evals.iter().cloned().fold(f64::INFINITY, f64::min);
let row_neg = evals.iter().filter(|&&v| v < neg_floor).count();
min_raw_eig = min_raw_eig.min(row_min);
max_raw_eig = max_raw_eig.max(max_abs);
if row_neg > 0 {
indefinite_rows += 1;
total_neg_dirs += row_neg;
}
}
let rel_min = min_raw_eig / max_raw_eig.max(1.0);
eprintln!(
"[#1190] real-data curvature anchor (K={k}, N={}): RAW assembled H_tt \
min_eig={min_raw_eig:.6e} (rel={rel_min:.3e}) indefinite_rows={indefinite_rows}/{} \
total_neg_dirs={total_neg_dirs}",
z_train.nrows(),
sys.rows.len()
);
assert!(
rel_min >= -1.0e-8,
"real-data curvature anchor is genuinely indefinite: raw assembled H_tt \
min eigenvalue {min_raw_eig:.6e} (relative {rel_min:.3e}) is negative on \
{indefinite_rows}/{} rows ({total_neg_dirs} negative directions) — the \
d=2 atoms are under-identified on real OLMo activations (#1190). The \
curvature anchor must be PD (or its negative directions must be genuine \
closed-form gauge nulls, not data-supported directions).",
sys.rows.len()
);
}
#[test]
pub(crate) fn olmo_real_arrival_floor_tracks_data_ceiling() {
let path = olmo_fixture_path("olmo_mixedlayer_pca64_768.npy");
let z = read_npy_f32_2d(&path);
assert_eq!(z.dim(), (768, 64), "real OLMo fixture shape");
let z_train = z.slice(s![..384, ..]).to_owned();
let k = 8usize;
let term = real_data_torus_seed_term(z_train.view(), k, 2);
let rho = SaeManifoldRho::new(0.0, 0.0, vec![array![0.0, 0.0]; k]);
let objective = SaeManifoldOuterObjective::new(
term,
z_train.clone(),
None,
rho.clone(),
0,
0.04,
1.0e-6,
1.0e-6,
);
let anchor = linear_span_anchor(&objective.term, z_train.view())
.expect("Eckart-Young anchor must be recoverable on the real fixture");
let sst = {
let mut means = vec![0.0_f64; z_train.ncols()];
for col in 0..z_train.ncols() {
let mut acc = 0.0;
for row in 0..z_train.nrows() {
acc += z_train[[row, col]];
}
means[col] = acc / z_train.nrows() as f64;
}
let mut s = 0.0_f64;
for row in 0..z_train.nrows() {
for col in 0..z_train.ncols() {
let c = z_train[[row, col]] - means[col];
s += c * c;
}
}
s
};
let anchor_ev = 1.0 - anchor.residual_norm_sq / sst;
assert!(
anchor_ev.is_finite() && anchor_ev > SAE_FIT_DATA_COLLAPSE_EV_FLOOR,
"real-data Eckart-Young anchor ceiling {anchor_ev:.5} is degenerate (#1189)."
);
eprintln!("[#1189] real-data anchor ceiling anchor_ev={anchor_ev:.5}");
let arrival_floor_k = |achievable_ceiling: f64, k_active: usize| -> f64 {
let k = k_active.max(1) as f64;
(achievable_ceiling * ((k - 1.0) / k)).max(SAE_FIT_DATA_COLLAPSE_EV_FLOOR)
};
let arrival_floor = |achievable_ceiling: f64| -> f64 { arrival_floor_k(achievable_ceiling, 1) };
let real_regime_ceiling = 0.40_f64; for k in [1usize, 2, 8] {
let f = arrival_floor_k(real_regime_ceiling, k);
eprintln!("[#1189] real regime K={k}: ceiling={real_regime_ceiling} floor={f:.5}");
assert!(
f < real_regime_ceiling,
"[#1189] arrival floor {f:.5} (K={k}) is not strictly below the achievable real-data \
ceiling {real_regime_ceiling}: a genuine fit AT the ceiling would be rejected and \
demoted to a structurally collapsed cascade."
);
}
let synthetic_ceiling = 0.95_f64;
for k in [1usize, 2, 8] {
let f = arrival_floor_k(synthetic_ceiling, k);
assert!(
f < synthetic_ceiling && 0.94 >= f,
"[#1189] synthetic floor {f:.5} (K={k}) must sit below the achievable ceiling \
{synthetic_ceiling} so a genuine planted-harmonic recovery (EV ≈ 0.94) is accepted."
);
}
let pathological_floor = arrival_floor(0.0);
assert!(
pathological_floor >= SAE_FIT_DATA_COLLAPSE_EV_FLOOR,
"the #1189 floor dropped below the data-collapse threshold on a pathological ceiling \
(floor {pathological_floor:.5} < {SAE_FIT_DATA_COLLAPSE_EV_FLOOR}) (#1189)."
);
for k in [1usize, 2, 8] {
let f = arrival_floor_k(anchor_ev, k);
assert!(
f >= SAE_FIT_DATA_COLLAPSE_EV_FLOOR
&& f < anchor_ev.max(SAE_FIT_DATA_COLLAPSE_EV_FLOOR + 1e-9),
"real-data anchor floor {f:.5} (K={k}) fell outside [{SAE_FIT_DATA_COLLAPSE_EV_FLOOR}, \
anchor ceiling {anchor_ev:.5}) (#1189)."
);
}
let k3_linear_ceiling = 0.30_f64;
let k3_curved_arrival = 0.2461_f64; let k3_floor = arrival_floor_k(k3_linear_ceiling, 3);
eprintln!(
"[#1026] K=3 ceiling={k3_linear_ceiling:.4} curved_arrival={k3_curved_arrival:.4} \
floor={k3_floor:.4}"
);
assert!(
k3_curved_arrival >= k3_floor,
"[#1026] the per-atom-share floor {k3_floor:.4} still demotes a genuine curved K=3 \
arrival at EV {k3_curved_arrival:.4} (linear ceiling {k3_linear_ceiling:.4}); the K>=2 \
co-collapse regression is NOT fixed."
);
assert!(
k3_floor < k3_linear_ceiling && k3_curved_arrival < k3_linear_ceiling,
"[#1026] the per-atom-share floor {k3_floor:.4} must sit strictly below the FULL linear \
ceiling {k3_linear_ceiling:.4} (else there is no forgiveness and the regression is \
vacuous), and the curved arrival {k3_curved_arrival:.4} must lie in that forgiven band."
);
let f1 = arrival_floor_k(k3_linear_ceiling, 1);
let f2 = arrival_floor_k(k3_linear_ceiling, 2);
let f3 = arrival_floor_k(k3_linear_ceiling, 3);
let f8 = arrival_floor_k(k3_linear_ceiling, 8);
assert!(
f1 <= f2 && f2 <= f3 && f3 <= f8,
"[#1026] arrival floor is not monotone non-decreasing across K \
(K=1 {f1:.4}, K=2 {f2:.4}, K=3 {f3:.4}, K=8 {f8:.4})."
);
assert!(
f8 < k3_linear_ceiling,
"[#1026] the share floor reached/exceeded the full ceiling at large K \
(K=8 {f8:.4} >= ceiling {k3_linear_ceiling:.4})."
);
}
pub(crate) fn olmo_fixture_path(name: &str) -> std::path::PathBuf {
let mani = std::path::Path::new(env!("CARGO_MANIFEST_DIR"));
let crate_local = mani.join("tests/data").join(name);
if crate_local.exists() {
return crate_local;
}
let workspace_root = mani.join("../../tests/data").join(name);
if workspace_root.exists() {
return workspace_root;
}
panic!(
"OLMo fixture {name} not found at {} or {}",
crate_local.display(),
workspace_root.display()
);
}
pub(crate) fn read_npy_f32_2d(path: &std::path::Path) -> Array2<f64> {
let bytes = std::fs::read(path).unwrap_or_else(|e| panic!("read {}: {e}", path.display()));
assert!(
bytes.len() > 10 && &bytes[0..6] == b"\x93NUMPY",
"not a .npy"
);
let major = bytes[6];
let (hdr_start, hdr_len) = if major == 1 {
(10usize, u16::from_le_bytes([bytes[8], bytes[9]]) as usize)
} else {
(
12usize,
u32::from_le_bytes([bytes[8], bytes[9], bytes[10], bytes[11]]) as usize,
)
};
let data_off = hdr_start + hdr_len;
let header = std::str::from_utf8(&bytes[hdr_start..data_off]).unwrap();
assert!(
header.contains("'<f4'") || header.contains("\"<f4\""),
"fixture must be little-endian float32; header: {header}"
);
assert!(!header.contains("True"), "fixture must be C-contiguous");
let open = header.find('(').unwrap();
let close = header[open..].find(')').unwrap() + open;
let dims: Vec<usize> = header[open + 1..close]
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(|s| s.parse::<usize>().unwrap())
.collect();
assert_eq!(dims.len(), 2, "fixture must be 2-D");
let (n, p) = (dims[0], dims[1]);
let mut out = Array2::<f64>::zeros((n, p));
let payload = &bytes[data_off..];
assert!(payload.len() >= n * p * 4, "truncated payload");
for r in 0..n {
for c in 0..p {
let i = (r * p + c) * 4;
let v =
f32::from_le_bytes([payload[i], payload[i + 1], payload[i + 2], payload[i + 3]]);
out[[r, c]] = v as f64;
}
}
out
}
#[test]
pub(crate) fn fit_data_collapse_verdict_uses_one_self_consistent_state_s1() {
let coords = array![[0.0_f64], [0.25], [0.5], [0.75]];
let n = coords.nrows();
let mut phi = Array2::<f64>::zeros((n, 3));
let mut jet = Array3::<f64>::zeros((n, 3, 1));
for row in 0..n {
let angle = 2.0 * std::f64::consts::PI * coords[[row, 0]];
phi[[row, 0]] = 1.0;
phi[[row, 1]] = angle.sin();
phi[[row, 2]] = angle.cos();
jet[[row, 1, 0]] = 2.0 * std::f64::consts::PI * angle.cos();
jet[[row, 2, 0]] = -2.0 * std::f64::consts::PI * angle.sin();
}
let mut live_decoder = Array2::<f64>::zeros((3, 2));
live_decoder[[2, 0]] = 1.0;
live_decoder[[1, 1]] = 1.0;
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"circle",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
live_decoder,
Array2::<f64>::eye(3),
)
.unwrap();
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
Array2::<f64>::zeros((n, 1)),
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::softmax(1.0),
)
.unwrap();
let mut live_term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
let target = array![[1.0, 0.0], [0.0, 1.0], [-1.0, 0.0], [0.0, -1.0]];
let rho = SaeManifoldRho::new(-0.3, 0.0, vec![array![0.0]]);
let live_verdict = live_term
.dictionary_collapse_verdict(target.view(), &rho, None)
.expect("live verdict");
assert!(!live_verdict.all_decoders_vanished(1));
assert!(!live_verdict.degenerate(1));
let recorded_live = live_term
.record_fit_data_collapse_if_needed(target.view(), &rho, 3)
.unwrap();
assert!(
!recorded_live,
"a live decoder must not produce a terminal collapse event"
);
assert!(
!live_term
.collapse_events()
.iter()
.any(|e| e.action == CollapseAction::Terminal),
"no terminal event may be recorded for a live decoder"
);
let mut vanished_term = live_term.clone();
vanished_term.atoms[0].decoder_coefficients_mut().fill(0.0);
let vanished_verdict = vanished_term
.dictionary_collapse_verdict(target.view(), &rho, None)
.expect("vanished verdict");
assert!(vanished_verdict.all_decoders_vanished(1));
assert!(vanished_verdict.degenerate(1));
let recorded_vanished = vanished_term
.record_fit_data_collapse_if_needed(target.view(), &rho, 7)
.unwrap();
assert!(
recorded_vanished,
"an exactly vanished decoder is a genuine #853/#976 co-collapse"
);
let terminal = vanished_term
.collapse_events()
.iter()
.find(|e| e.action == CollapseAction::Terminal)
.expect("a terminal collapse event must be recorded for the vanished dictionary");
assert!(
terminal.floor.is_finite() && terminal.floor >= 0.0,
"the event carries the derived signal boundary, never a NaN sentinel"
);
}
#[test]
pub(crate) fn fast_encode_matches_per_row_warm_start() {
let path = olmo_fixture_path("olmo_l18_pca64_635.npy");
let z = read_npy_f32_2d(&path);
let n = z.nrows();
let k = 1usize;
let term = real_data_torus_seed_term(z.view(), k, 3);
let mut norm_bound = 0.0_f64;
for r in 0..n {
norm_bound = norm_bound.max(z.row(r).dot(&z.row(r)).sqrt());
}
let atlas = crate::encode::EncodeAtlas::build(
&term.atoms,
&vec![1.0_f64; k],
norm_bound,
crate::encode::AtlasConfig::default(),
)
.expect("atlas builds");
let atom = &term.atoms[0];
let amps = ndarray::Array1::<f64>::ones(n);
let mut ref_coords = ndarray::Array2::<f64>::zeros((n, atom.latent_dim()));
let mut ref_valid = vec![false; n];
for row in 0..n {
if let Some((cidx, _)) =
crate::encode::nearest_chart(&atlas.atoms[0], z.row(row), amps[row])
{
if let Some(t) = crate::encode::amortized_warm_start(
&atlas.atoms[0].charts[cidx],
z.row(row),
amps[row],
) {
ref_coords.row_mut(row).assign(&t);
ref_valid[row] = true;
}
}
}
let (fast_coords, fast_valid) = atlas
.amortized_encode_batch_fast(atom, 0, z.view(), amps.view())
.expect("batched fast encode runs");
let mut max_diff = 0.0_f64;
for row in 0..n {
assert_eq!(
fast_valid[row], ref_valid[row],
"valid-mask mismatch at row {row} (routing/predictor disagreement)"
);
if ref_valid[row] {
for c in 0..atom.latent_dim() {
max_diff = max_diff.max((fast_coords[[row, c]] - ref_coords[[row, c]]).abs());
}
}
}
assert!(
max_diff < 1.0e-12,
"batched fast-encode must match the per-row warm-start to 1e-12 (same affine \
map, GEMM-batched); max|Δcoord| = {max_diff:.3e}"
);
assert!(
ref_valid.iter().filter(|&&v| v).count() > n / 2,
"fixture must produce valid encodes on most rows; got {}",
ref_valid.iter().filter(|&&v| v).count()
);
}
#[test]
pub(crate) fn fast_reconstruct_matches_per_row_decode() {
let path = olmo_fixture_path("olmo_l18_pca64_635.npy");
let z = read_npy_f32_2d(&path);
let n = z.nrows();
let p = z.ncols();
let k = 1usize;
let term = real_data_torus_seed_term(z.view(), k, 3);
let mut norm_bound = 0.0_f64;
for r in 0..n {
norm_bound = norm_bound.max(z.row(r).dot(&z.row(r)).sqrt());
}
let atlas = crate::encode::EncodeAtlas::build(
&term.atoms,
&vec![1.0_f64; k],
norm_bound,
crate::encode::AtlasConfig::default(),
)
.expect("atlas builds");
let atom = &term.atoms[0];
let amps = ndarray::Array1::<f64>::ones(n);
let evaluator = atom.basis_evaluator.as_ref().unwrap().clone();
let (fast_recon, fast_valid) = atlas
.amortized_reconstruct_batch_fast(atom, 0, z.view(), amps.view())
.expect("batched fast reconstruct runs");
let (coords, enc_valid) = atlas
.amortized_encode_batch_fast(atom, 0, z.view(), amps.view())
.expect("batched fast encode runs");
let mut max_diff = 0.0_f64;
let mut valid_rows = 0usize;
for row in 0..n {
assert_eq!(
fast_valid[row], enc_valid[row],
"reconstruct valid-mask must equal encode valid-mask at row {row}"
);
if !fast_valid[row] {
for col in 0..p {
assert_eq!(
fast_recon[[row, col]],
0.0,
"uncertified row {row} must decode to zero, got {}",
fast_recon[[row, col]]
);
}
continue;
}
valid_rows += 1;
let single = coords.row(row).insert_axis(ndarray::Axis(0)).to_owned();
let (phi_row, _jet) = evaluator
.evaluate(single.view())
.expect("single basis eval");
let decoded_row = phi_row.dot(atom.decoder_coefficients()); for col in 0..p {
let expect = amps[row] * decoded_row[[0, col]];
max_diff = max_diff.max((fast_recon[[row, col]] - expect).abs());
}
}
assert!(
max_diff < 1.0e-10,
"batched fast reconstruct must match the per-row decode z·Φ(t̂)·B (same GEMM, \
batched basis eval); max|Δrecon| = {max_diff:.3e}"
);
assert!(
valid_rows > n / 2,
"fixture must reconstruct most rows; got {valid_rows} valid of {n}"
);
}
#[test]
fn fast_forward_is_accuracy_parity_with_certified() {
let (z_tr, z) = olmo_l18_oos_split();
let n = z.nrows();
let p = z.ncols();
let term = real_data_torus_seed_term(z_tr.view(), 1, 3);
let mut norm_bound = 0.0_f64;
for r in 0..z_tr.nrows() {
norm_bound = norm_bound.max(z_tr.row(r).dot(&z_tr.row(r)).sqrt());
}
for r in 0..n {
norm_bound = norm_bound.max(z.row(r).dot(&z.row(r)).sqrt());
}
let atlas = crate::encode::EncodeAtlas::build(
&term.atoms,
&vec![1.0_f64; 1],
norm_bound,
crate::encode::AtlasConfig::default(),
)
.unwrap();
let atom = &term.atoms[0];
let amps = ndarray::Array1::<f64>::ones(n);
let evaluator = atom.basis_evaluator.as_ref().unwrap().clone();
let (fast_recon, fast_valid) = atlas
.amortized_reconstruct_batch_fast(atom, 0, z.view(), amps.view())
.unwrap();
let mut both: Vec<(f64, f64)> = Vec::new(); let mut fast_valid_count = 0usize;
let mut cert_valid_count = 0usize;
for row in 0..n {
let xr = z.row(row);
let xn = xr.dot(&xr).sqrt().max(1e-12);
let fast_e = if fast_valid[row] {
fast_valid_count += 1;
let mut e = 0.0;
for c in 0..p {
let d = fast_recon[[row, c]] - xr[c];
e += d * d;
}
Some(e.sqrt() / xn)
} else {
None
};
let (coords, cert) = atlas.certified_encode_row(atom, 0, xr, amps[row]).unwrap();
let cert_e = if cert.beta.is_finite() && cert.h.is_finite() {
cert_valid_count += 1;
let single = coords.insert_axis(ndarray::Axis(0));
let (phi, _) = evaluator.evaluate(single.view()).unwrap();
let dec = phi.dot(atom.decoder_coefficients());
let mut e = 0.0;
for c in 0..p {
let d = amps[row] * dec[[0, c]] - xr[c];
e += d * d;
}
Some(e.sqrt() / xn)
} else {
None
};
if let (Some(f), Some(c)) = (fast_e, cert_e) {
both.push((f, c));
}
}
assert!(
fast_valid_count >= cert_valid_count,
"fast path must cover >= certified rows; fast={fast_valid_count} cert={cert_valid_count}"
);
assert!(
both.len() > n / 8,
"need a non-trivial co-valid set; got {} of {n}",
both.len()
);
let mean = |v: &[f64]| v.iter().sum::<f64>() / v.len().max(1) as f64;
let fast_mean = mean(&both.iter().map(|x| x.0).collect::<Vec<_>>());
let cert_mean = mean(&both.iter().map(|x| x.1).collect::<Vec<_>>());
assert!(
fast_mean <= 1.05 * cert_mean,
"fast forward must be accuracy-parity with certified on co-valid rows; \
fast_mean={fast_mean:.4} cert_mean={cert_mean:.4} ratio={:.3}",
fast_mean / cert_mean
);
}
pub(crate) fn olmo_l18_oos_split() -> (Array2<f64>, Array2<f64>) {
let path = olmo_fixture_path("olmo_l18_pca64_635.npy");
let z = read_npy_f32_2d(&path);
let n = z.nrows();
let n_tr = (n * 6) / 10;
(
z.slice(s![..n_tr, ..]).to_owned(),
z.slice(s![n_tr.., ..]).to_owned(),
)
}
fn oos_sq_sum(z: &Array2<f64>) -> f64 {
let mut t = 0.0;
for r in 0..z.nrows() {
for c in 0..z.ncols() {
t += z[[r, c]] * z[[r, c]];
}
}
t
}
pub(crate) fn oos_train_curved(
z_tr: &Array2<f64>,
z_te: &Array2<f64>,
d: usize,
h: usize,
iters: usize,
data_driven: bool,
maxc: usize,
) -> (f64, f64) {
let n_tr = z_tr.nrows();
let n_te = z_te.nrows();
let p = z_tr.ncols();
let tot_tr = oos_sq_sum(z_tr);
let tot_te = oos_sq_sum(z_te);
let mut nb = 0.0_f64;
for r in 0..n_tr {
nb = nb.max(z_tr.row(r).dot(&z_tr.row(r)).sqrt());
}
for r in 0..n_te {
nb = nb.max(z_te.row(r).dot(&z_te.row(r)).sqrt());
}
let ev_eval = Arc::new(TorusHarmonicEvaluator::new(d, h).unwrap());
let seed =
sae_pca_seed_initial_coords(z_tr.view(), &vec![SaeAtomBasisKind::Periodic; 1], &vec![d])
.unwrap();
let mut coords = seed.slice(s![0, .., 0..d]).to_owned();
let build = |coords: &Array2<f64>| -> SaeManifoldAtom {
let (phi, jet) = ev_eval.evaluate(coords.view()).unwrap();
let m = phi.ncols();
let mut xtx = fast_ata(&phi);
for i in 0..m {
xtx[[i, i]] += 1e-8;
}
let xtz = fast_atb(&phi, &z_tr.to_owned());
let dec = xtx.cholesky(Side::Lower).unwrap().solve_mat(&xtz);
SaeManifoldAtom::new_with_provided_function_gram(
"t",
SaeAtomBasisKind::Periodic,
d,
phi,
jet,
dec,
Array2::eye(m),
)
.unwrap()
.with_basis_evaluator(ev_eval.clone())
};
let mk_atlas = |atom: &SaeManifoldAtom, coords: &Array2<f64>| {
if data_driven {
crate::encode::EncodeAtlas::build_data_driven(
std::slice::from_ref(atom),
std::slice::from_ref(coords),
&[1.0],
nb,
maxc,
crate::encode::AtlasConfig::default(),
)
.unwrap()
} else {
crate::encode::EncodeAtlas::build(
std::slice::from_ref(atom),
&[1.0],
nb,
crate::encode::AtlasConfig::default(),
)
.unwrap()
}
};
let amps_tr = ndarray::Array1::<f64>::ones(n_tr);
let mut atom = build(&coords);
for _ in 0..iters {
let atlas = mk_atlas(&atom, &coords);
let (ec, v) = atlas
.amortized_encode_batch_fast(&atom, 0, z_tr.view(), amps_tr.view())
.unwrap();
for i in 0..n_tr {
if v[i] {
coords.row_mut(i).assign(&ec.row(i));
}
}
atom = build(&coords);
}
let rt = atom.basis_values.dot(atom.decoder_coefficients());
let mut etr = 0.0;
for r in 0..n_tr {
for c in 0..p {
let dd = rt[[r, c]] - z_tr[[r, c]];
etr += dd * dd;
}
}
let ev_in = 1.0 - etr / tot_tr;
let atlas = mk_atlas(&atom, &coords);
let amps_te = ndarray::Array1::<f64>::ones(n_te);
let (rte, _vm) = atlas
.amortized_reconstruct_batch_fast(&atom, 0, z_te.view(), amps_te.view())
.unwrap();
let mut ete = 0.0;
for r in 0..n_te {
for c in 0..p {
let dd = rte[[r, c]] - z_te[[r, c]];
ete += dd * dd;
}
}
let ev_oos = 1.0 - ete / tot_te;
(ev_in, ev_oos)
}
#[test]
fn curved_atom_oos_competitive_with_real_topk_sae() {
let (tr, te) = olmo_l18_oos_split();
let (_in, curved_oos) = oos_train_curved(&tr, &te, 2, 3, 5, false, 0);
eprintln!(
"curved d=2 OOS EV={curved_oos:.4} (real TopK SAE k=2 OOS ≈ 0.217–0.242, \
see tests/sae/real_topk_sae_baseline.py)"
);
assert!(
curved_oos > 0.15 && curved_oos < 0.40,
"curved d=2 OOS EV must sit in the real-TopK-SAE-competitive band [0.15,0.40] \
(measured ~0.22); got {curved_oos:.4}"
);
}
#[test]
fn more_harmonics_overfit_out_of_sample() {
let (tr, te) = olmo_l18_oos_split();
let (in3, oos3) = oos_train_curved(&tr, &te, 2, 3, 5, false, 0);
let (in4, oos4) = oos_train_curved(&tr, &te, 2, 4, 5, false, 0);
eprintln!("h=3: in={in3:.4} OOS={oos3:.4} h=4: in={in4:.4} OOS={oos4:.4}");
assert!(
in4 - in3 > 0.03,
"extra harmonic must raise IN-SAMPLE EV (capacity added); in3={in3:.4} in4={in4:.4}"
);
assert!(
oos4 - oos3 < 0.01,
"extra harmonic must NOT improve OOS EV (it overfits); oos3={oos3:.4} oos4={oos4:.4}"
);
}
#[test]
fn manifold_training_loop_generalizes_but_overfits_out_of_sample() {
let (tr, te) = olmo_l18_oos_split();
let (in0, oos0) = oos_train_curved(&tr, &te, 2, 3, 0, false, 0); let (in6, oos6) = oos_train_curved(&tr, &te, 2, 3, 6, false, 0); eprintln!("seed: in={in0:.4} OOS={oos0:.4} trained: in={in6:.4} OOS={oos6:.4}");
assert!(
oos6 > oos0,
"training must improve OOS reconstruction over the seed; oos0={oos0:.4} oos6={oos6:.4}"
);
assert!(
in6 - in0 > 3.0 * (oos6 - oos0),
"in-sample gain must dwarf OOS gain (overfitting); din={:.4} doos={:.4}",
in6 - in0,
oos6 - oos0
);
}
#[test]
fn data_driven_higher_latent_dim_helps_out_of_sample() {
let (tr, te) = olmo_l18_oos_split();
let (_in2, oos_d2) = oos_train_curved(&tr, &te, 2, 1, 5, true, 256);
let (_in4, oos_d4) = oos_train_curved(&tr, &te, 4, 1, 5, true, 256);
eprintln!("OOS data-driven d=2 EV={oos_d2:.4} d=4 EV={oos_d4:.4}");
assert!(
oos_d2 > 0.0 && oos_d4 > 1.3 * oos_d2,
"data-driven d=4 must beat d=2 OUT-OF-SAMPLE by >30% (latent-dim unlock \
generalises); oos_d2={oos_d2:.4} oos_d4={oos_d4:.4}"
);
}
fn oos_linear_affine_rank_ev(z_tr: &Array2<f64>, z_te: &Array2<f64>, r: usize) -> f64 {
use gam_linalg::faer_ndarray::FaerEigh;
let p = z_tr.ncols();
let n_tr = z_tr.nrows();
let mut mean = ndarray::Array1::<f64>::zeros(p);
for row in 0..n_tr {
for c in 0..p {
mean[c] += z_tr[[row, c]];
}
}
mean.mapv_inplace(|v| v / n_tr as f64);
let mut centered_tr = z_tr.clone();
for row in 0..n_tr {
for c in 0..p {
centered_tr[[row, c]] -= mean[c];
}
}
let cov = fast_ata(¢ered_tr);
let (_evals, evecs) = cov.eigh(faer::Side::Lower).unwrap();
let r = r.min(p);
let mut v = Array2::<f64>::zeros((p, r));
for j in 0..r {
let col = p - 1 - j;
for i in 0..p {
v[[i, j]] = evecs[[i, col]];
}
}
let n_te = z_te.nrows();
let mut err = 0.0_f64;
let mut tot = 0.0_f64;
for row in 0..n_te {
let mut coords = ndarray::Array1::<f64>::zeros(r);
for j in 0..r {
let mut acc = 0.0_f64;
for c in 0..p {
acc += (z_te[[row, c]] - mean[c]) * v[[c, j]];
}
coords[j] = acc;
}
for c in 0..p {
let mut recon = mean[c];
for j in 0..r {
recon += coords[j] * v[[c, j]];
}
let d = z_te[[row, c]] - recon;
err += d * d;
tot += z_te[[row, c]] * z_te[[row, c]];
}
}
1.0 - err / tot
}
fn planted_curve_oos_split() -> (Array2<f64>, Array2<f64>) {
const P: usize = 6;
let amp = [1.0, 1.0, 0.7, 0.7, 0.5, 0.5];
let embed = |theta: f64| -> [f64; P] {
[
amp[0] * theta.cos(),
amp[1] * theta.sin(),
amp[2] * (2.0 * theta).cos(),
amp[3] * (2.0 * theta).sin(),
amp[4] * (3.0 * theta).cos(),
amp[5] * (3.0 * theta).sin(),
]
};
let two_pi = std::f64::consts::TAU;
let build = |n: usize, phase: f64, seed0: u64| -> Array2<f64> {
let mut z = Array2::<f64>::zeros((n, P));
let mut state = seed0;
for row in 0..n {
let theta = two_pi * ((row as f64 + phase) / n as f64);
let x = embed(theta);
for c in 0..P {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let unit = ((state >> 11) as f64) * f64::from_bits(0x3CA0000000000000);
z[[row, c]] = x[c] + 0.02 * (2.0 * unit - 1.0);
}
}
z
};
(
build(300, 0.0, 0x9E3779B97F4A7C15),
build(200, 0.5, 0xD1B54A32D192ED03),
)
}
#[test]
fn curved_warm_start_matches_or_beats_linear_baseline_out_of_sample_2261() {
let (tr, te) = planted_curve_oos_split();
let linear_oos = oos_linear_affine_rank_ev(&tr, &te, 2);
let (_seed_in, seed_oos) = oos_train_curved(&tr, &te, 1, 3, 0, false, 0);
let (_fit_in, fit_oos) = oos_train_curved(&tr, &te, 1, 3, 5, false, 0);
eprintln!(
"[#2261] planted curve OOS EV: linear rank-2={linear_oos:.4} \
curved seed(0 iters)={seed_oos:.4} curved fit(5 iters)={fit_oos:.4}"
);
assert!(
seed_oos >= linear_oos - 1.0e-3,
"PCA-warm-started curved seed must start AT OR ABOVE the linear rank-2 \
baseline (that is what warm-starting from the PCA/linear solution buys); \
seed_oos={seed_oos:.4} linear_oos={linear_oos:.4}"
);
assert!(
fit_oos > linear_oos + 0.05,
"curved arm must BEAT the closed-form linear rank-2 baseline out-of-sample \
on genuinely curved data (its higher harmonics are invisible to any linear \
code); fit_oos={fit_oos:.4} linear_oos={linear_oos:.4}"
);
assert!(
linear_oos < 0.85,
"planted curve must be genuinely nonlinear (rank-2 linear OOS EV must be \
well below 1); linear_oos={linear_oos:.4}"
);
}
fn production_circle_coords_at_seed(
target: &Array2<f64>,
random_state: u64,
) -> ndarray::Array1<f64> {
use crate::manifold::{
SaeFitAssignmentKind, SaeFitConfig, SaeFitRequest, SaeFitSeedReport, SaeFitSeedRequest,
SaeMinimalSeedReport, SaeMinimalSeedRequest, build_sae_fit_seed, build_sae_minimal_seed,
run_sae_manifold_fit,
};
let assignment_kind = SaeFitAssignmentKind::Softmax;
let minimal = build_sae_minimal_seed(SaeMinimalSeedRequest {
target: target.view(),
atom_basis: vec!["periodic".to_string()],
atom_dim: vec![1],
assignment_kind,
alpha: 1.0,
tau: 1.0,
threshold: 0.0,
top_k: None,
random_state,
initial_logits: None,
initial_coords: None,
})
.expect("minimal seed");
let SaeMinimalSeedReport {
geometry_plans,
basis_values,
basis_jacobian,
decoder_coefficients,
smooth_penalties,
initial_logits,
initial_coords,
refine_routing,
} = minimal;
let registry = AnalyticPenaltyRegistry::new();
let seed = build_sae_fit_seed(SaeFitSeedRequest {
target: target.view(),
geometry_plans: &geometry_plans,
basis_values: basis_values.view(),
basis_jacobian: basis_jacobian.view(),
decoder_coefficients: decoder_coefficients.view(),
smooth_penalties: smooth_penalties.view(),
initial_logits: initial_logits.view(),
initial_coords: initial_coords.view(),
alpha: 1.0,
tau: 1.0,
learnable_alpha: false,
assignment_kind,
sparsity_strength: 1.0,
smoothness: 1.0,
max_iter: 12,
learning_rate: 1.0,
ridge_ext_coord: 1.0e-6,
ridge_beta: 1.0e-6,
top_k: None,
threshold: 0.0,
native_ard_enabled: true,
seed_refine_routing: refine_routing,
seed_refine_random_state: random_state,
data_row_reseed: false,
fit_config: SaeFitConfig::default(),
temperature_schedule: None,
fisher_metric: None,
row_loss_weights: None,
registry: ®istry,
})
.expect("fit seed");
let SaeFitSeedReport {
base_term,
initial_rho,
isometry_pin_active,
metric_provenance,
} = seed;
let report = run_sae_manifold_fit(SaeFitRequest {
reconstruction_optimism_folds: None,
base_term,
target: target.clone(),
registry,
initial_rho,
max_iter: 12,
learning_rate: 1.0,
ridge_ext_coord: 1.0e-6,
ridge_beta: 1.0e-6,
alpha: 1.0,
isometry_pin_active,
metric_provenance,
promote_from_residual: false,
run_structure_search: false,
run_outer_rho_search: false,
structured_residual_passes: 0,
cancel: None,
})
.expect("production circle fit")
.manifold_or_error()
.expect("planted circle must retain a manifold atom");
let coords = report.term.assignment.coords[0].as_matrix();
ndarray::Array1::from_iter((0..coords.nrows()).map(|i| coords[[i, 0]]))
}
#[test]
fn production_circle_readout_cross_seed_concordance_2260() {
let path = olmo_fixture_path("qwen35_9b_actsL21_pca64_2000.npy");
let full = read_npy_f32_2d(&path);
let n = 800.min(full.nrows());
let z = full.slice(s![..n, ..]).to_owned();
let seeds = [42u64, 43, 44, 45, 46];
let mut coord_rows = Array2::<f64>::zeros((seeds.len(), n));
for (r, &seed) in seeds.iter().enumerate() {
let coords = production_circle_coords_at_seed(&z, seed);
assert_eq!(coords.len(), n, "seed {seed}: one coordinate per row");
assert!(
coords.iter().all(|v| v.is_finite()),
"seed {seed}: converged circle coordinate must be finite"
);
coord_rows.row_mut(r).assign(&coords);
}
let report = crate::circular_concordance::circular_concordance(coord_rows.view(), 1.0)
.expect("circular concordance over the five seed replicates");
let min_aligned = report.minimum_aligned_score;
let mean_aligned = report.mean_aligned_score;
eprintln!(
"[#2260] Qwen-9B L21 production circle readout, seeds 42-46 (N={n}): \
cross-seed circular concordance min={min_aligned:?} mean={mean_aligned:?} \
(torch-lane reported 0.67-0.97 spread; 1.0 = seed-identical ordering)"
);
for pair in &report.pairs {
eprintln!(
"[#2260] pair ({},{}) aligned={:?} reflected={:?}",
pair.left, pair.right, pair.aligned_score, pair.reflected
);
}
assert!(
report.coverage.iter().all(|c| c.well_posed),
"every seed's circle embedding must be well-posed (2-D span) for the \
concordance to be meaningful"
);
assert!(
min_aligned.is_some() && mean_aligned.is_some(),
"cross-seed aligned concordance must be computable across the five seeds"
);
let min_aligned = min_aligned.expect("min aligned concordance");
assert!(
min_aligned >= 0.99,
"production circle readout must be seed-stable: min cross-seed aligned \
concordance {min_aligned:.4} must be >= 0.99 (deterministic atan2 seed); \
a lower value means seed-dependent basin selection has regressed (#2260)"
);
}
#[test]
fn certified_encode_is_globally_sound_near_self_crossing() {
use ndarray::{Array1, Array2};
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(5).unwrap());
let n_seed = 64usize;
let seed: Array2<f64> = Array2::from_shape_fn((n_seed, 1), |(i, _)| i as f64 / n_seed as f64);
let (phi, jet) = evaluator.evaluate(seed.view()).unwrap();
let m = phi.ncols();
let mut decoder = Array2::<f64>::zeros((m, 2));
decoder[[2, 0]] = 1.0; decoder[[3, 1]] = 1.0; let atom = SaeManifoldAtom::new_with_provided_function_gram(
"fig8",
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(m),
)
.unwrap()
.with_basis_evaluator(evaluator.clone());
let recon = |t: f64| -> [f64; 2] {
let a = 2.0 * std::f64::consts::PI * t;
[a.cos(), (2.0 * a).sin()]
};
let grad = |t: f64, x: &[f64; 2]| -> f64 {
let a = 2.0 * std::f64::consts::PI * t;
let dm = [
-2.0 * std::f64::consts::PI * a.sin(),
4.0 * std::f64::consts::PI * (2.0 * a).cos(),
];
let r = recon(t);
-(dm[0] * (x[0] - r[0]) + dm[1] * (x[1] - r[1]))
};
let global_min_err = |x: &[f64; 2]| -> f64 {
let mut best = f64::INFINITY;
let g = 20000;
for i in 0..g {
let t = i as f64 / g as f64;
let r = recon(t);
let e = (r[0] - x[0]).powi(2) + (r[1] - x[1]).powi(2);
if e < best {
best = e;
}
}
best.sqrt()
};
let atlas = crate::encode::EncodeAtlas::build(
std::slice::from_ref(&atom),
&[1.0],
1.6,
crate::encode::AtlasConfig {
grid_resolution: 64,
ridge: 1e-10,
newton_steps: 8,
},
)
.unwrap();
let mut certified = 0usize;
let mut worst_grad = 0.0_f64;
let mut worst_global_excess = 0.0_f64;
let steps = 41;
for ix in 0..steps {
for iy in 0..steps {
let x0 = -0.30 + 0.60 * ix as f64 / (steps - 1) as f64;
let x1 = -0.30 + 0.60 * iy as f64 / (steps - 1) as f64;
let xv = Array1::from(vec![x0, x1]);
let (coord, cert) = atlas
.certified_encode_row(&atom, 0, xv.view(), 1.0)
.unwrap();
if !cert.certified() {
continue;
}
certified += 1;
let t = coord[0];
worst_grad = worst_grad.max(grad(t, &[x0, x1]).abs());
let r = recon(t);
let cert_err = ((r[0] - x0).powi(2) + (r[1] - x1).powi(2)).sqrt();
worst_global_excess = worst_global_excess.max(cert_err - global_min_err(&[x0, x1]));
}
}
eprintln!(
"certified={certified}/{} worst|grad|={worst_grad:.2e} worst global excess={worst_global_excess:.4}",
steps * steps
);
assert!(
worst_grad < 1e-4,
"certificate's LOCAL claim must hold: certified coords must be stationary \
points (‖∇‖≈0); worst |∇| = {worst_grad:.2e}"
);
assert!(
certified > steps * steps / 2,
"fixture must certify most targets; got {certified}"
);
assert!(
worst_global_excess < 5e-3,
"certified encode must be GLOBALLY sound (top-K routing): worst excess over \
the global min = {worst_global_excess:.5} (was ~0.08 with single-chart routing)"
);
}