#[cfg(test)]
mod tests {
use crate::encode::{AtlasConfig, EncodeAtlas};
use crate::inference::steering::set_coordinate;
use crate::manifold::{
SaeFisherRowMetricRequest, SaeFitAssignmentKind, SaeFitConfig, SaeFitSeedReport,
SaeFitSeedRequest, SaeManifoldTerm, SaeMinimalSeedReport, SaeMinimalSeedRequest,
build_sae_fit_seed, build_sae_minimal_seed,
};
use gam_terms::analytic_penalties::AnalyticPenaltyRegistry;
use ndarray::{Array2, Array3};
const N_MONTHS: usize = 12;
const N_ROWS: usize = 96;
const P_OUT: usize = 6;
const NOISE_SIGMA: f64 = 0.01;
fn lcg(state: &mut u64) -> f64 {
*state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*state >> 11) as f64) / ((1u64 << 53) as f64)
}
fn lcg_normal(state: &mut u64) -> f64 {
let u1 = lcg(state).max(1e-12);
let u2 = lcg(state);
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
fn month_embedding_target() -> Array2<f64> {
let mut state = 0x2263_0000_0000_0003u64;
Array2::from_shape_fn((N_ROWS, P_OUT), |(i, j)| {
let theta = std::f64::consts::TAU * (i as f64) / (N_ROWS as f64);
let harmonic = (j / 2 + 1) as f64;
let clean = if j % 2 == 0 {
(harmonic * theta).cos()
} else {
(harmonic * theta).sin()
};
clean + NOISE_SIGMA * lcg_normal(&mut state)
})
}
fn build_month_term() -> (SaeManifoldTerm, Array2<f64>) {
let target = month_embedding_target();
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: 0,
initial_logits: None,
initial_coords: None,
})
.expect("minimal seed on the month circle");
let SaeMinimalSeedReport {
geometry_plans,
basis_values,
basis_jacobian,
decoder_coefficients,
smooth_penalties,
initial_logits,
initial_coords,
refine_routing,
} = minimal;
let identity_u = Array3::<f64>::from_shape_fn(
(N_ROWS, P_OUT, 1),
|(_, i, _)| if i == 0 { 1.0 } else { 0.0 },
);
let metric_request = SaeFisherRowMetricRequest::from_tag(
identity_u.view(),
N_ROWS,
P_OUT,
None,
Some("uncertified_approximation"),
None,
)
.expect("behavioral metric request");
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: 40,
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: 0,
data_row_reseed: false,
fit_config: SaeFitConfig::default(),
temperature_schedule: None,
fisher_metric: Some(metric_request),
row_loss_weights: None,
registry: ®istry,
})
.expect("fit seed on the month circle");
let SaeFitSeedReport {
base_term: term, ..
} = seed;
(term, target)
}
#[test]
fn fixture_chart_spans_the_circle() {
let (term, _) = build_month_term();
let coords = term.assignment.coords[0].as_matrix();
let mut lo = f64::INFINITY;
let mut hi = f64::NEG_INFINITY;
for row in 0..coords.nrows() {
lo = lo.min(coords[[row, 0]]);
hi = hi.max(coords[[row, 0]]);
}
let span_turns = hi - lo;
let mut distinct: Vec<f64> = (0..coords.nrows()).map(|r| coords[[r, 0]]).collect();
distinct.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
distinct.dedup_by(|a, b| (*a - *b).abs() <= 1e-6);
eprintln!(
"[#2263 fixture] fitted chart spans {span_turns:.4} turns over {} rows, \
{} distinct coordinates",
coords.nrows(),
distinct.len()
);
assert!(
span_turns >= 0.8,
"the fitted chart spans only {span_turns:.4} turns; a displacement measured on it \
says nothing about steering"
);
assert!(
distinct.len() >= N_MONTHS,
"the fitted chart collapsed to {} distinct coordinates for {} rows; it cannot \
resolve {N_MONTHS} months",
distinct.len(),
coords.nrows()
);
}
fn months_between(from: f64, to: f64) -> f64 {
let turns = to - from;
let mut months = turns * N_MONTHS as f64;
while months <= -(N_MONTHS as f64) / 2.0 {
months += N_MONTHS as f64;
}
while months > N_MONTHS as f64 / 2.0 {
months -= N_MONTHS as f64;
}
months
}
fn atlas_for(term: &SaeManifoldTerm, target: &Array2<f64>, amplitude: f64) -> EncodeAtlas {
let mut norm_bound = 0.0_f64;
for row in 0..target.nrows() {
norm_bound = norm_bound.max(target.row(row).dot(&target.row(row)).sqrt());
}
EncodeAtlas::build(
&term.atoms,
&[amplitude.max(1.0)],
norm_bound * (1.0 + 2.0 * amplitude.max(1.0)),
AtlasConfig::default(),
)
.expect("encode atlas builds over the month circle")
}
#[test]
fn requested_month_advance_is_realized_exactly_at_the_rows_own_amplitude() {
let (term, target) = build_month_term();
let metric = term.row_metric().expect("metric installed").clone();
let atlas = atlas_for(&term, &target, 1.0);
let atom = &term.atoms[0];
let coords = term.assignment.coords[0].as_matrix();
let mut worst = 0.0_f64;
let mut worst_case = (0usize, 0i32, 0.0_f64);
for requested in 1..=6i32 {
let mut realized_by_row = Vec::new();
for row in 0..N_ROWS {
let (t_from_read, _) = atlas
.certified_encode_row(atom, 0, target.row(row), 1.0)
.expect("the unedited row encodes");
let t_from = t_from_read[0];
let encode_gap = months_between(coords[[row, 0]], t_from);
if requested == 1 {
eprintln!(
"[#2263 encode-control] row {row}: fitted coord {:+.6} vs certified \
encode {t_from:+.6} (gap {encode_gap:+.4} months)",
coords[[row, 0]]
);
}
let t_to = t_from + (requested as f64) / (N_MONTHS as f64);
let set = set_coordinate(
&term,
&metric,
&atlas,
target.row(row),
0,
row,
1.0,
&[t_to],
)
.expect("set_coordinate writes the requested month");
let (t_realized, _) = atlas
.certified_encode_row(atom, 0, set.edited.view(), 1.0)
.expect("the edited row re-encodes");
let realized = months_between(set.t_from_certified[0], t_realized[0]);
realized_by_row.push(realized);
let antipodal = requested * 2 == N_MONTHS as i32;
let error = if antipodal {
(realized.abs() - requested as f64).abs()
} else {
(realized - requested as f64).abs()
};
if error > worst {
worst = error;
worst_case = (row, requested, realized);
}
}
let mean =
realized_by_row.iter().sum::<f64>() / realized_by_row.len() as f64;
eprintln!(
"[#2263 displacement] requested +{requested} months: realized mean {mean:+.4} \
over {} rows (min {:+.4}, max {:+.4})",
realized_by_row.len(),
realized_by_row.iter().cloned().fold(f64::INFINITY, f64::min),
realized_by_row
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max),
);
}
eprintln!(
"[#2263 displacement] worst |realized − requested| = {worst:.3e} months \
(row {}, requested +{}, realized {:+.4})",
worst_case.0, worst_case.1, worst_case.2
);
assert!(
worst <= 0.15,
"the chart round trip does not realize the requested month advance: \
worst |realized − requested| = {worst:.3e} months (row {}, requested +{}, \
realized {:+.4})",
worst_case.0,
worst_case.1,
worst_case.2
);
}
#[test]
fn requests_past_a_full_turn_wrap_rather_than_saturate() {
let (term, target) = build_month_term();
let metric = term.row_metric().expect("metric installed").clone();
let atlas = atlas_for(&term, &target, 1.0);
let atom = &term.atoms[0];
let mut worst = 0.0_f64;
for requested in 7..=18i32 {
let mut realized_by_row = Vec::new();
for row in 0..N_ROWS {
let Ok((t_from_read, _)) =
atlas.certified_encode_row(atom, 0, target.row(row), 1.0)
else {
continue;
};
let t_to = t_from_read[0] + (requested as f64) / (N_MONTHS as f64);
let Ok(set) = set_coordinate(
&term, &metric, &atlas, target.row(row), 0, row, 1.0, &[t_to],
) else {
continue;
};
let Ok((t_realized, _)) =
atlas.certified_encode_row(atom, 0, set.edited.view(), 1.0)
else {
continue;
};
realized_by_row
.push(months_between(set.t_from_certified[0], t_realized[0]));
}
if realized_by_row.is_empty() {
continue;
}
let mut expected = requested as f64;
while expected > N_MONTHS as f64 / 2.0 {
expected -= N_MONTHS as f64;
}
let antipodal = (expected.abs() - N_MONTHS as f64 / 2.0).abs() < 1e-9;
let mean =
realized_by_row.iter().sum::<f64>() / realized_by_row.len() as f64;
let row_worst = realized_by_row.iter().fold(0.0_f64, |m, &r| {
let e = if antipodal {
(r.abs() - expected.abs()).abs()
} else {
(r - expected).abs()
};
m.max(e)
});
eprintln!(
"[#2263 wrap] requested +{requested}: wrap predicts {expected:+.0}, realized mean {mean:+.4}, worst row error {row_worst:.4} months"
);
worst = worst.max(row_worst);
}
assert!(
worst <= 0.15,
"requests past a full turn neither wrap nor track: worst deviation from the \
wrap prediction is {worst:.4} months. A saturating gain would fail here while \
+1..+5 passed"
);
}
#[test]
fn realized_advance_scales_with_the_written_amplitude_not_the_request() {
let (term, target) = build_month_term();
let metric = term.row_metric().expect("metric installed").clone();
let atom = &term.atoms[0];
let requested = 1i32;
for &alpha in &[1.0_f64, 2.0, 4.0, 8.0, 16.0] {
let atlas = atlas_for(&term, &target, alpha);
let mut realized_by_row = Vec::new();
for row in 0..N_ROWS {
let Ok((t_from_read, _)) =
atlas.certified_encode_row(atom, 0, target.row(row), alpha)
else {
continue;
};
let t_to = t_from_read[0] + (requested as f64) / (N_MONTHS as f64);
let Ok(set) = set_coordinate(
&term, &metric, &atlas, target.row(row), 0, row, alpha, &[t_to],
) else {
continue;
};
let Ok((t_realized, _)) =
atlas.certified_encode_row(atom, 0, set.edited.view(), 1.0)
else {
continue;
};
realized_by_row.push(months_between(
set.t_from_certified[0],
t_realized[0],
));
}
if realized_by_row.is_empty() {
eprintln!(
"[#2263 amplitude] alpha={alpha}: no row produced a certified \
encode of the edited activation"
);
continue;
}
let mean =
realized_by_row.iter().sum::<f64>() / realized_by_row.len() as f64;
let mut sorted = realized_by_row.clone();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
eprintln!(
"[#2263 amplitude] requested +{requested} month at alpha={alpha}: \
realized mean {mean:+.4}, median {:+.4}, over {} of {N_ROWS} rows",
sorted[sorted.len() / 2],
sorted.len(),
);
}
}
}