use ndarray::{Array2, Array3};
use super::atom::{SaeAtomBasisKind, SaeManifoldAtom};
use super::inframe_curved::{
ChartOccupancyStatus, CurvedRegion, InFrameCurvedConfig, WeightFrameOccupancy,
activate_residual_frame, dense_ambient_radial_reference, fit_inframe_curved_regions,
fit_inframe_curved_weight_frame_catalog, inframe_curved_region_prediction, residual_span_frame,
};
use super::weight_frame_catalog::{
WeightFrameCatalogConfig, WeightFrameMatrix, WeightFrameSource,
frame_catalog_from_weight_matrices,
};
struct Lcg(u64);
impl Lcg {
fn new(seed: u64) -> Self {
Self(seed | 1)
}
fn next_unit(&mut self) -> f64 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let bits = (self.0 >> 11) as f64 / (1u64 << 53) as f64; 2.0 * bits - 1.0
}
fn normal(&mut self) -> f64 {
let u1 = (self.next_unit() * 0.5 + 0.5).max(1.0e-12);
let u2 = self.next_unit() * 0.5 + 0.5;
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
}
fn random_orthonormal(p: usize, r: usize, seed: u64) -> Array2<f64> {
let mut rng = Lcg::new(seed);
let mut q = Array2::<f64>::zeros((p, r));
for col in 0..r {
let mut v: Vec<f64> = (0..p).map(|_| rng.normal()).collect();
for prev in 0..col {
let mut dot = 0.0;
for i in 0..p {
dot += v[i] * q[[i, prev]];
}
for i in 0..p {
v[i] -= dot * q[[i, prev]];
}
}
let mut norm = 0.0;
for i in 0..p {
norm += v[i] * v[i];
}
norm = norm.sqrt().max(1.0e-12);
for i in 0..p {
q[[i, col]] = v[i] / norm;
}
}
q
}
fn planted_curved_residual(
n: usize,
p: usize,
r_true: usize,
shell_noise: f64,
ambient_noise: f64,
seed: u64,
) -> (Array2<f64>, Array2<f64>) {
let q = random_orthonormal(p, r_true, seed);
let mut rng = Lcg::new(seed ^ 0x9E3779B97F4A7C15);
let mut latent = Array2::<f64>::zeros((n, r_true));
for i in 0..n {
let mut v: Vec<f64> = (0..r_true).map(|_| rng.normal()).collect();
let mut norm = 0.0;
for x in &v {
norm += x * x;
}
norm = norm.sqrt().max(1.0e-12);
let radius = 1.0 + shell_noise * rng.normal();
for x in &mut v {
*x = radius * *x / norm;
}
for j in 0..r_true {
latent[[i, j]] = v[j];
}
}
let mut residual = Array2::<f64>::zeros((n, p));
for i in 0..n {
for j in 0..p {
let mut acc = 0.0;
for k in 0..r_true {
acc += latent[[i, k]] * q[[j, k]];
}
residual[[i, j]] = acc + ambient_noise * rng.normal();
}
}
(residual, q)
}
fn project_matrix_onto_frame(matrix: &Array2<f64>, frame: &Array2<f64>) -> Array2<f64> {
frame.dot(&frame.t().dot(matrix))
}
fn relative_frobenius(a: &Array2<f64>, b: &Array2<f64>) -> f64 {
let mut diff = 0.0;
let mut denom = 0.0;
for (x, y) in a.iter().zip(b.iter()) {
let d = x - y;
diff += d * d;
denom += y * y;
}
(diff / denom.max(1.0e-30)).sqrt()
}
#[test]
fn weight_frame_catalog_spans_component_column_images() {
let p = 24;
let q_ov = random_orthonormal(p, 2, 1001);
let q_mlp = random_orthonormal(p, 2, 1002);
let mut w_v = Array2::<f64>::zeros((2, p));
for j in 0..p {
w_v[[0, j]] = 0.7 * (j as f64 + 1.0);
w_v[[1, j]] = if j % 2 == 0 { 1.0 } else { -0.5 };
}
let ov = q_ov.dot(&w_v);
let mut down_coeff = Array2::<f64>::zeros((2, 5));
for j in 0..5 {
down_coeff[[0, j]] = 1.0 + j as f64;
down_coeff[[1, j]] = if j % 2 == 0 { 2.0 } else { -1.0 };
}
let w_down = q_mlp.dot(&down_coeff);
let components = vec![
WeightFrameMatrix::attention_head_ov(3, 14, q_ov.view(), w_v.view()).expect("OV builds"),
WeightFrameMatrix::mlp_down_projection(3, w_down.view()),
];
let catalog = frame_catalog_from_weight_matrices(
&components,
&WeightFrameCatalogConfig {
frame_rank_min: 2,
frame_rank_max: 2,
..Default::default()
},
)
.expect("catalog builds");
assert_eq!(catalog.entries().len(), 2);
assert_eq!(
catalog.entries()[0].source,
WeightFrameSource::AttentionHeadOv { layer: 3, head: 14 }
);
assert_eq!(
catalog.entries()[1].source,
WeightFrameSource::MlpDownProjection { layer: 3 }
);
let ov_frame = catalog.entries()[0].frame.frame().to_owned();
let mlp_frame = catalog.entries()[1].frame.frame().to_owned();
let ov_projected = project_matrix_onto_frame(&ov, &ov_frame);
let mlp_projected = project_matrix_onto_frame(&w_down, &mlp_frame);
assert!(
relative_frobenius(&ov_projected, &ov) < 1.0e-10,
"OV catalog frame must span exactly the OV column image"
);
assert!(
relative_frobenius(&mlp_projected, &w_down) < 1.0e-10,
"MLP catalog frame must span exactly the W_down column image"
);
}
#[test]
fn weight_sourced_atom_fit_is_tagged_with_component_source() {
let n = 240;
let p = 96;
let r = 4;
let (residual, q) = planted_curved_residual(n, p, r, 0.02, 0.0, 1414);
let mut w_v = Array2::<f64>::zeros((r, p));
for j in 0..p {
for k in 0..r {
w_v[[k, j]] = ((j + 1 + k) as f64).sin();
}
}
let components =
vec![WeightFrameMatrix::attention_head_ov(2, 14, q.view(), w_v.view()).expect("OV")];
let catalog = frame_catalog_from_weight_matrices(
&components,
&WeightFrameCatalogConfig {
frame_rank_min: r,
frame_rank_max: r,
..Default::default()
},
)
.expect("catalog");
let config = InFrameCurvedConfig {
frame_rank_min: r,
frame_rank_max: r,
min_rows: 16,
..Default::default()
};
let result = fit_inframe_curved_weight_frame_catalog(
residual.view(),
&catalog,
&[WeightFrameOccupancy {
frame_index: 0,
rows: (0..n).collect(),
basis_size: 5,
}],
n,
&config,
)
.expect("weight-frame fit");
assert_eq!(result.records.len(), 1);
assert_eq!(
result.records[0].occupancy_status,
ChartOccupancyStatus::Occupied
);
assert_eq!(
result.records[0].frame_source,
Some(WeightFrameSource::AttentionHeadOv { layer: 2, head: 14 }),
"occupied atom record must carry native mechanism attribution"
);
if !result.curved_prediction.regions().is_empty() {
assert_eq!(
result.curved_prediction.regions()[0].frame_source(),
Some(&WeightFrameSource::AttentionHeadOv { layer: 2, head: 14 })
);
}
}
#[test]
fn zero_occupancy_weight_frame_is_reported_chartable_unoccupied() {
let n = 160;
let p = 64;
let r = 3;
let (residual, q_used) = planted_curved_residual(n, p, r, 0.03, 0.0, 5150);
let q_unused = random_orthonormal(p, r, 5151);
let mut coeff = Array2::<f64>::zeros((r, p));
for j in 0..p {
for k in 0..r {
coeff[[k, j]] = ((j + 2 * k + 1) as f64).cos();
}
}
let components = vec![
WeightFrameMatrix::attention_head_ov(6, 1, q_used.view(), coeff.view()).expect("OV"),
WeightFrameMatrix::mlp_down_projection(6, q_unused.view()),
];
let catalog = frame_catalog_from_weight_matrices(
&components,
&WeightFrameCatalogConfig {
frame_rank_min: r,
frame_rank_max: r,
..Default::default()
},
)
.expect("catalog");
let config = InFrameCurvedConfig {
frame_rank_min: r,
frame_rank_max: r,
min_rows: 16,
..Default::default()
};
let result = fit_inframe_curved_weight_frame_catalog(
residual.view(),
&catalog,
&[
WeightFrameOccupancy {
frame_index: 0,
rows: (0..n).collect(),
basis_size: 4,
},
WeightFrameOccupancy {
frame_index: 1,
rows: Vec::new(),
basis_size: 4,
},
],
n,
&config,
)
.expect("weight-frame atlas");
assert_eq!(result.records.len(), 2);
let unoccupied = &result.records[1];
assert_eq!(
unoccupied.occupancy_status,
ChartOccupancyStatus::ChartableUnoccupied,
"unused weight frame remains in the atlas as chartable but unoccupied here"
);
assert_eq!(
unoccupied.frame_source,
Some(WeightFrameSource::MlpDownProjection { layer: 6 })
);
assert_eq!(unoccupied.evidence.n_rows, 0);
assert!(!unoccupied.evidence.selected_by_bic);
}
#[test]
fn planted_low_rank_curved_recovered_inframe_p2048() {
let n = 1200;
let p = 2048;
let r_true = 6;
let (residual, _q) = planted_curved_residual(n, p, r_true, 0.02, 0.0, 42);
let m = 8usize;
let config = InFrameCurvedConfig {
frame_rank_min: 2,
frame_rank_max: 16,
crossfit_folds: 4,
min_rows: 32,
..Default::default()
};
let region = CurvedRegion {
rows: (0..n).collect(),
basis_size: m,
};
let start = std::time::Instant::now();
let result = fit_inframe_curved_regions(residual.view(), &[region], n, &config)
.expect("in-frame curved fit");
let elapsed = start.elapsed();
let rec = &result.records[0];
assert!(
rec.frame_rank >= r_true && rec.frame_rank <= 16,
"frame rank {} should recover the intrinsic rank {r_true} within the band",
rec.frame_rank
);
assert!(
rec.frame_rank < p,
"frame rank {} must be far below ambient p={p}",
rec.frame_rank
);
assert!(
rec.evidence.deviance_gain > 0.0,
"curved chart should improve held-out deviance, got {}",
rec.evidence.deviance_gain
);
assert_eq!(result.selected_regions, vec![0], "planted region selected");
let ledger = &result.ledger;
assert_eq!(ledger.dense_border_coeffs, m * p);
assert_eq!(ledger.inframe_border_coeffs, m * rec.frame_rank);
let border_shrink = ledger.border_shrink();
let cov_shrink = ledger.cov_shrink();
assert!(
border_shrink >= (p as f64 / 16.0),
"border must shrink by ~p/r; got {border_shrink:.1}x (dense={} inframe={})",
ledger.dense_border_coeffs,
ledger.inframe_border_coeffs
);
assert!(
cov_shrink >= 10_000.0,
"posterior covariance must shrink by (p/r)²; got {cov_shrink:.0}x"
);
assert!(
ledger.dense_cov_bytes >= 1_000_000_000,
"dense (M·p)² covariance should be ~GB, got {} bytes",
ledger.dense_cov_bytes
);
assert!(
ledger.inframe_cov_bytes <= 4_000_000,
"in-frame (M·r)² covariance should be well under a MB, got {} bytes",
ledger.inframe_cov_bytes
);
assert!(
elapsed.as_secs_f64() < 30.0,
"in-frame fit at p={p} took {elapsed:?}; should be fast"
);
assert_eq!(result.curved_prediction.n_rows(), n);
assert_eq!(result.curved_prediction.output_dim(), p);
assert_eq!(
result.curved_prediction.inframe_entries(),
n * rec.frame_rank,
"hot prediction storage must be the accepted region's N_g x r image"
);
assert_eq!(
result.curved_prediction.accepted_ambient_entries_if_eager(),
n * p,
"the eager ambient atom image would have been N_g x p"
);
assert!(
result.curved_prediction.inframe_entries()
< result.curved_prediction.accepted_ambient_entries_if_eager(),
"curved prediction must stay in-frame on the hot path"
);
let pred = result.curved_prediction.materialize_ambient();
let mut mean_norm = 0.0;
for i in 0..n {
let mut ss = 0.0;
for j in 0..p {
ss += pred[[i, j]] * pred[[i, j]];
}
mean_norm += ss.sqrt();
}
mean_norm /= n as f64;
assert!(
(mean_norm - 1.0).abs() < 0.25,
"reconstructed shell radius {mean_norm:.3} should be ~1"
);
}
#[test]
fn inframe_matches_dense_full_p_when_frame_contains_truth() {
let n = 300;
let p = 256;
let r_true = 4;
let (residual, _q) = planted_curved_residual(n, p, r_true, 0.1, 0.0, 7);
let config = InFrameCurvedConfig {
frame_rank_min: r_true,
frame_rank_max: r_true,
min_rows: 16,
..Default::default()
};
let rows: Vec<usize> = (0..n).collect();
let inframe_pred = inframe_curved_region_prediction(residual.view(), &rows, &config)
.expect("in-frame prediction")
.expect("frame learned");
let rank = inframe_pred.frame_rank();
assert_eq!(rank, r_true, "frame rank pinned to the true intrinsic rank");
assert_eq!(
inframe_pred.inframe_entries(),
n * r_true,
"single-region prediction stays in the learned r-frame"
);
let dense_pred =
dense_ambient_radial_reference(residual.view(), config.whitening_ridge).expect("dense fit");
let inframe_ambient = inframe_pred.materialize_ambient();
let mut diff = 0.0;
let mut denom = 0.0;
for i in 0..n {
for j in 0..p {
let d = inframe_ambient[[i, j]] - dense_pred[[i, j]];
diff += d * d;
denom += dense_pred[[i, j]] * dense_pred[[i, j]];
}
}
let rel = (diff / denom.max(1.0e-30)).sqrt();
assert!(
rel < 1.0e-6,
"in-frame and dense full-p radial fits must match on the frame's span; \
relative Frobenius diff {rel:.3e}"
);
}
#[test]
fn residual_span_frame_is_the_production_hook_low_rank_and_spans_truth() {
let n = 400;
let p = 1024;
let r_true = 6;
let (residual, _q) = planted_curved_residual(n, p, r_true, 0.05, 0.0, 2130);
let config = InFrameCurvedConfig {
frame_rank_min: r_true,
frame_rank_max: 16,
min_rows: 16,
..Default::default()
};
let rows: Vec<usize> = (0..n).collect();
let frame = residual_span_frame(residual.view(), &rows, &config)
.expect("frame learns")
.expect("beneficial low-rank frame exists for a planted low-rank residual");
let r = frame.rank();
assert!(
r >= r_true && r <= 16 && r < p,
"seam frame rank {r} should recover the intrinsic rank {r_true} and stay far below p={p}"
);
let u = frame.frame().to_owned(); let mut diff = 0.0;
let mut denom = 0.0;
for i in 0..n {
let mut z = vec![0.0; r];
for (k, zk) in z.iter_mut().enumerate() {
let mut acc = 0.0;
for j in 0..p {
acc += residual[[i, j]] * u[[j, k]];
}
*zk = acc;
}
for j in 0..p {
let mut lifted = 0.0;
for (k, &zk) in z.iter().enumerate() {
lifted += zk * u[[j, k]];
}
let d = residual[[i, j]] - lifted;
diff += d * d;
denom += residual[[i, j]] * residual[[i, j]];
}
}
let rel = (diff / denom.max(1e-30)).sqrt();
assert!(
rel < 1e-6,
"seam frame must span the planted subspace (residual reconstructs through U Uᵀ); rel={rel:.3e}"
);
let mut rng = Lcg::new(9001);
let mut iso = Array2::<f64>::zeros((64, 8));
for i in 0..64 {
for j in 0..8 {
iso[[i, j]] = rng.normal();
}
}
let tight = InFrameCurvedConfig {
frame_rank_min: 2,
frame_rank_max: 4,
rank_cutoff: 1e-9, ..Default::default()
};
let iso_rows: Vec<usize> = (0..64).collect();
let got = residual_span_frame(iso.view(), &iso_rows, &tight).expect("runs");
if let Some(f) = got {
assert!(
f.rank() <= 4 && f.rank() < 8,
"seam frame must stay strictly low-rank"
);
}
}
#[test]
fn activate_residual_frame_installs_factored_decoder_and_engages_frames() {
let n = 200;
let p = 128;
let m = 3usize; let r_true = 5;
let (residual, _q) = planted_curved_residual(n, p, r_true, 0.05, 0.0, 4242);
let mut rng = Lcg::new(77);
let mut decoder = Array2::<f64>::zeros((m, p));
for a in 0..m {
for j in 0..p {
decoder[[a, j]] = rng.normal();
}
}
let basis_values = Array2::<f64>::zeros((1, m));
let basis_jacobian = Array3::<f64>::zeros((1, m, 1));
let smooth_penalty = Array2::<f64>::eye(m);
let mut atom = SaeManifoldAtom::new(
"seam",
SaeAtomBasisKind::Periodic,
1,
basis_values,
basis_jacobian,
decoder,
smooth_penalty,
)
.expect("atom builds");
assert!(atom.decoder_frame.is_none(), "starts on the full-p path");
let config = InFrameCurvedConfig {
frame_rank_min: r_true,
frame_rank_max: 16,
min_rows: 16,
..Default::default()
};
let rows: Vec<usize> = (0..n).collect();
let r = activate_residual_frame(&mut atom, residual.view(), &rows, &config)
.expect("activation runs")
.expect("beneficial low-rank frame installed");
assert!(
r >= r_true && r < p,
"installed frame rank {r} low-rank vs p={p}"
);
let frame = atom.decoder_frame.as_ref().expect("frame installed");
assert_eq!(frame.rank(), r);
let u = frame.frame().to_owned();
let mut reproj = atom.decoder_coefficients.dot(&u).dot(&u.t());
reproj -= &atom.decoder_coefficients;
let mut fro = 0.0;
for v in reproj.iter() {
fro += v * v;
}
assert!(
fro.sqrt() < 1e-9,
"activated decoder must satisfy B = (B U) Uᵀ exactly; residual {:.3e}",
fro.sqrt()
);
}
#[test]
fn linear_structure_is_not_promoted_to_curved() {
let n = 800;
let p = 512;
let dir = random_orthonormal(p, 1, 314);
let mut rng = Lcg::new(2130);
let mut residual = Array2::<f64>::zeros((n, p));
for i in 0..n {
let s = rng.normal(); for j in 0..p {
residual[[i, j]] = s * dir[[j, 0]];
}
}
let config = InFrameCurvedConfig {
frame_rank_min: 4,
frame_rank_max: 8,
min_rows: 32,
..Default::default()
};
let region = CurvedRegion {
rows: (0..n).collect(),
basis_size: 8,
};
let result = fit_inframe_curved_regions(residual.view(), &[region], n, &config)
.expect("fit runs on linear structure");
assert!(
result.selected_regions.is_empty(),
"purely linear (rank-1) structure must NOT be promoted to a curved atom; \
deviance_gain={} margin={}",
result.records[0].evidence.deviance_gain,
result.records[0].evidence.margin
);
}
#[test]
fn ledger_shrink_matches_reviewer_frontier_shape() {
let n = 600;
let p = 4096;
let r_target = 16;
let (residual, _q) = planted_curved_residual(n, p, r_target, 0.02, 0.0, 99);
let config = InFrameCurvedConfig {
frame_rank_min: r_target,
frame_rank_max: r_target,
min_rows: 32,
..Default::default()
};
let region = CurvedRegion {
rows: (0..n).collect(),
basis_size: 8,
};
let result = fit_inframe_curved_regions(residual.view(), &[region], n, &config).expect("fit");
assert_eq!(result.records[0].frame_rank, r_target);
if result.selected_regions.is_empty() {
let rec = &result.records[0];
assert_eq!(rec.dense_border_coeffs, 8 * p);
assert_eq!(rec.inframe_border_coeffs, 8 * r_target);
} else {
let ledger = &result.ledger;
assert_eq!(ledger.dense_border_coeffs, 8 * p);
assert_eq!(ledger.inframe_border_coeffs, 8 * r_target);
assert_eq!(ledger.dense_cov_bytes, (8 * p) * (8 * p) * 8);
assert_eq!(
ledger.inframe_cov_bytes,
(8 * r_target) * (8 * r_target) * 8
);
assert_eq!(ledger.inframe_cov_bytes, 131_072);
assert!((ledger.border_shrink() - 256.0).abs() < 1.0e-9);
assert!((ledger.cov_shrink() - 65_536.0).abs() < 1.0e-6);
}
}
#[test]
fn inframe_curved_p4096_feasible_where_dense_joint_ooms_2134() {
let n = 1500; let p = 4096; let r_true = 8;
let m = 8usize; let (residual, _q) = planted_curved_residual(n, p, r_true, 0.02, 0.0, 21_34);
let config = InFrameCurvedConfig {
frame_rank_min: 2,
frame_rank_max: 16,
crossfit_folds: 4,
min_rows: 32,
..Default::default()
};
let region = CurvedRegion {
rows: (0..n).collect(),
basis_size: m,
};
let result = fit_inframe_curved_regions(residual.view(), &[region], n, &config)
.expect("in-frame curved fit is feasible at p=4096 where the dense joint OOMs");
let rec = &result.records[0];
assert!(
rec.frame_rank >= r_true && rec.frame_rank <= 16 && rec.frame_rank < p,
"frame rank {} recovers intrinsic rank {r_true} and stays far below p={p}",
rec.frame_rank
);
assert_eq!(
result.selected_regions,
vec![0],
"planted curved region selected"
);
let ledger = &result.ledger;
assert_eq!(ledger.dense_border_coeffs, m * p);
assert_eq!(ledger.inframe_border_coeffs, m * rec.frame_rank);
assert!(
ledger.dense_cov_bytes >= 8_000_000_000,
"dense (M·p)² covariance is the ~8.6 GB the joint lane OOMs on, got {} bytes",
ledger.dense_cov_bytes
);
assert!(
ledger.inframe_cov_bytes <= 1_000_000,
"in-frame (M·r)² covariance must stay well under a MB, got {} bytes",
ledger.inframe_cov_bytes
);
eprintln!(
"[#2134 p4096] N={n} p={p} frame_rank={} dense_border={} inframe_border={} \
dense_cov_bytes={} inframe_cov_bytes={} border_shrink={:.1} cov_shrink={:.1}",
rec.frame_rank,
ledger.dense_border_coeffs,
ledger.inframe_border_coeffs,
ledger.dense_cov_bytes,
ledger.inframe_cov_bytes,
ledger.border_shrink(),
ledger.cov_shrink(),
);
assert_eq!(
result.curved_prediction.inframe_entries(),
n * rec.frame_rank
);
assert!(
result.curved_prediction.inframe_entries()
< result.curved_prediction.accepted_ambient_entries_if_eager(),
"curved prediction must stay in-frame (N_g×r), never the eager N_g×p ambient image"
);
}
#[test]
fn accepted_curved_prediction_hot_path_stays_in_r_frame() {
let n = 180;
let p = 128;
let r_true = 4;
let (residual, _q) = planted_curved_residual(n, p, r_true, 0.02, 0.0, 974);
let config = InFrameCurvedConfig {
frame_rank_min: r_true,
frame_rank_max: r_true,
min_rows: 16,
..Default::default()
};
let region = CurvedRegion {
rows: (0..n).collect(),
basis_size: 5,
};
let result = fit_inframe_curved_regions(residual.view(), &[region], n, &config).expect("fit");
assert_eq!(
result.selected_regions,
vec![0],
"planted curved region selected"
);
let prediction = &result.curved_prediction;
assert_eq!(prediction.regions().len(), 1);
assert_eq!(prediction.regions()[0].frame_rank(), r_true);
assert_eq!(
prediction.inframe_entries(),
n * r_true,
"accepted atom image must be stored as N_g x r"
);
assert_eq!(
prediction.accepted_ambient_entries_if_eager(),
n * p,
"the forbidden eager atom image would be N_g x p"
);
assert!(
prediction.inframe_entries() * (p / r_true)
<= prediction.accepted_ambient_entries_if_eager(),
"in-frame storage should scale by r instead of p"
);
let slice = prediction.materialize_rows(&[0, n / 2, n - 1]);
assert_eq!(slice.dim(), (3, p));
let mut slice_energy = 0.0;
for value in slice.iter() {
slice_energy += value * value;
}
assert!(
slice_energy > 0.0,
"ambient lifting happens only for the requested residual slice"
);
}