#[cfg(test)]
mod tests {
use crate::inference::steering::{
AppliedDoseObservation, SteerPlan, TargetDoseConfig, TargetDoseRequest, steer_delta,
steer_to_target_nats,
};
use crate::manifold::{
SaeFisherRowMetricRequest, SaeFitAssignmentKind, SaeFitConfig, SaeFitSeedReport,
SaeFitSeedRequest, SaeManifoldTerm, SaeMinimalSeedReport, SaeMinimalSeedRequest,
build_sae_fit_seed, build_sae_minimal_seed,
};
use gam_problem::{FisherFactorKind, RowMetric};
use gam_terms::analytic_penalties::AnalyticPenaltyRegistry;
use ndarray::{Array1, Array2, Array3, ArrayView1};
use std::sync::Arc;
const N_CIRCLE: usize = 48;
const P_OUT: usize = 6;
const NOISE_SIGMA: f64 = 0.02;
const TIER0_SCALE: [f64; P_OUT] = [0.4, 2.5, 0.6, 1.8, 0.5, 2.2];
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 softmax(z: ArrayView1<'_, f64>) -> Vec<f64> {
let max_z = z.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let exps: Vec<f64> = z.iter().map(|&v| (v - max_z).exp()).collect();
let sum: f64 = exps.iter().sum();
exps.iter().map(|&v| v / sum).collect()
}
fn kl(p: &[f64], q: &[f64]) -> f64 {
p.iter()
.zip(q.iter())
.map(|(&pi, &qi)| if pi > 0.0 { pi * (pi / qi).ln() } else { 0.0 })
.sum()
}
fn categorical_quad_form(p_probs: &[f64], delta: &[f64]) -> f64 {
let mut s1 = 0.0_f64;
let mut s2 = 0.0_f64;
for (&pi, &di) in p_probs.iter().zip(delta.iter()) {
s1 += pi * di * di;
s2 += pi * di;
}
s1 - s2 * s2
}
fn categorical_fisher_metric(z_raw: &Array2<f64>) -> RowMetric {
let n = z_raw.nrows();
let p = z_raw.ncols();
let mut flat = vec![0.0_f64; n * p * p];
for row in 0..n {
let probs = softmax(z_raw.row(row));
for c in 0..p {
let sqrt_pc = probs[c].sqrt();
for i in 0..p {
let e_ci = if i == c { 1.0 } else { 0.0 };
flat[row * p * p + i * p + c] = sqrt_pc * (e_ci - probs[i]);
}
}
}
let u = Array2::from_shape_vec((n, p * p), flat).expect("U shape");
RowMetric::output_fisher(Arc::new(u), p, p)
.expect("full-rank output-Fisher metric")
.with_fisher_factor_kind(FisherFactorKind::ExactFull)
.expect("closed-form categorical factor is exact")
}
fn circle_embedding_target() -> Array2<f64> {
let mut state = 0x2249_0000_0000_0011u64;
Array2::from_shape_fn((N_CIRCLE, P_OUT), |(i, j)| {
let theta = std::f64::consts::TAU * (i as f64) / (N_CIRCLE 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_calibrated_term() -> (SaeManifoldTerm, RowMetric, Array1<f64>) {
let target = circle_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");
let SaeMinimalSeedReport {
geometry_plans,
basis_values,
basis_jacobian,
decoder_coefficients,
smooth_penalties,
initial_logits,
initial_coords,
refine_routing,
} = minimal;
let dummy_u =
Array3::<f64>::from_shape_fn(
(N_CIRCLE, P_OUT, 1),
|(_, i, _)| if i == 0 { 1.0 } else { 0.0 },
);
let dummy_metric = SaeFisherRowMetricRequest::from_tag(
dummy_u.view(),
N_CIRCLE,
P_OUT,
None,
Some("uncertified_approximation"),
None,
)
.expect("placeholder 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: 4,
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(dummy_metric),
row_loss_weights: None,
registry: ®istry,
})
.expect("fit seed");
let SaeFitSeedReport {
base_term: mut term,
..
} = seed;
let scale = Array1::from_vec(TIER0_SCALE.to_vec());
term.set_tier0_scale(scale.clone())
.expect("inject asymmetric tier0 scale");
let coords = term.assignment.coords[0].as_matrix();
let atom = &term.atoms[0];
let mut z_raw = Array2::<f64>::zeros((N_CIRCLE, P_OUT));
for row in 0..N_CIRCLE {
let t = coords.row(row).to_vec();
let t_mat = Array2::from_shape_vec((1, t.len()), t).expect("coord row");
let g_int = atom
.decode_at_coords(t_mat.view())
.expect("decode at fitted coord");
for j in 0..P_OUT {
z_raw[[row, j]] = scale[j] * g_int[[0, j]];
}
}
let metric = categorical_fisher_metric(&z_raw);
term.set_row_metric(metric)
.expect("install calibrated metric");
let metric = term.row_metric().expect("metric installed").clone();
let angles: Array1<f64> = Array1::from_shape_fn(N_CIRCLE, |i| coords[[i, 0]]);
(term, metric, angles)
}
fn regress(xs: &[f64], ys: &[f64]) -> (f64, f64, f64) {
let n = xs.len() as f64;
let mean_x = xs.iter().sum::<f64>() / n;
let mean_y = ys.iter().sum::<f64>() / n;
let mut sxx = 0.0_f64;
let mut sxy = 0.0_f64;
let mut syy = 0.0_f64;
for (&x, &y) in xs.iter().zip(ys.iter()) {
sxx += (x - mean_x) * (x - mean_x);
sxy += (x - mean_x) * (y - mean_y);
syy += (y - mean_y) * (y - mean_y);
}
let slope = sxy / sxx;
let intercept = mean_y - slope * mean_x;
let mut ss_res = 0.0_f64;
for (&x, &y) in xs.iter().zip(ys.iter()) {
let pred = slope * x + intercept;
ss_res += (y - pred) * (y - pred);
}
let r2 = 1.0 - ss_res / syy;
(slope, intercept, r2)
}
#[test]
fn shipped_predicted_nats_calibrates_through_tier0_frame() {
let (term, metric, angles) = build_calibrated_term();
let atom = &term.atoms[0];
let tau = std::f64::consts::TAU;
let mut predicted = Vec::new();
let mut true_kl = Vec::new();
let mut internal_dose = Vec::new();
let scale = term.tier0_scale().expect("scale set").to_owned();
let mut state = 0x2249_C0FFEE_u64;
for row in 0..N_CIRCLE {
let t_from = angles[row];
let z_from: Vec<f64> = {
let t_mat = Array2::from_shape_vec((1, 1), vec![t_from]).unwrap();
let g = atom.decode_at_coords(t_mat.view()).unwrap();
(0..P_OUT).map(|j| scale[j] * g[[0, j]]).collect()
};
let p_from = softmax(ArrayView1::from(&z_from));
for _ in 0..6 {
let dt = (lcg(&mut state) - 0.5) * 0.02;
let t_to = (t_from + dt).rem_euclid(tau);
let chord: Vec<f64> = {
let t_mat = Array2::from_shape_vec((1, 1), vec![t_to]).unwrap();
let g = atom.decode_at_coords(t_mat.view()).unwrap();
(0..P_OUT)
.map(|j| scale[j] * g[[0, j]] - z_from[j])
.collect()
};
let chord_norm = chord.iter().map(|&c| c * c).sum::<f64>().sqrt();
if !(chord_norm > 0.0) {
continue;
}
let target_norm = 0.002 + 0.018 * lcg(&mut state);
let amp = target_norm / chord_norm;
let plan = steer_delta(&term, &metric, 0, row, amp, &[t_from], &[t_to])
.expect("steer_delta on calibrated term");
let pred = plan.predicted_nats.expect("behavioral dose");
let delta_raw = plan.delta.to_vec();
let z_to: Vec<f64> = z_from
.iter()
.zip(delta_raw.iter())
.map(|(&z, &d)| z + d)
.collect();
let p_to = softmax(ArrayView1::from(&z_to));
let exact = kl(&p_from, &p_to);
let delta_int: Vec<f64> = delta_raw
.iter()
.zip(scale.iter())
.map(|(&d, &s)| d / s)
.collect();
let mispriced = 0.5 * categorical_quad_form(&p_from, &delta_int);
let closed = 0.5 * categorical_quad_form(&p_from, &delta_raw);
assert!(
(pred - closed).abs() <= 1e-9 * closed.max(1e-30) + 1e-12,
"shipped predicted_nats {pred} must equal the closed-form categorical \
quadratic {closed} on the raw chord"
);
predicted.push(pred);
true_kl.push(exact);
internal_dose.push(mispriced);
}
}
let (slope, intercept, r2) = regress(&predicted, &true_kl);
let max_true = true_kl.iter().cloned().fold(0.0_f64, f64::max);
assert!(
slope > 0.97 && slope < 1.03,
"in-frame calibration slope {slope} must be ≈1 (predicted_nats tracks true KL)"
);
assert!(
intercept.abs() < 0.02 * max_true,
"in-frame calibration intercept {intercept} must be ≈0 vs max KL {max_true}"
);
assert!(
r2 > 0.995,
"in-frame calibration R² {r2} must be ≈1 across edit directions"
);
let (slope_int, _, _) = regress(&internal_dose, &true_kl);
assert!(
(slope_int - 1.0).abs() > 0.1,
"dropped-σ (internal-frame) dose slope {slope_int} must be materially off 1 under \
the asymmetric Tier-0 scale — the frame un-scaling is load-bearing, not cosmetic"
);
}
#[test]
fn target_dose_loop_lands_requested_nats() {
let (term, metric, angles) = build_calibrated_term();
let atom = &term.atoms[0];
let tau = std::f64::consts::TAU;
let scale = term.tier0_scale().expect("scale set").to_owned();
let row = 3usize;
let t_from = angles[row];
let t_to = (t_from + 0.02).rem_euclid(tau);
let z_from: Vec<f64> = {
let t_mat = Array2::from_shape_vec((1, 1), vec![t_from]).unwrap();
let g = atom.decode_at_coords(t_mat.view()).unwrap();
(0..P_OUT).map(|j| scale[j] * g[[0, j]]).collect()
};
let p_from = softmax(ArrayView1::from(&z_from));
let dg_raw: Vec<f64> = {
let t_mat = Array2::from_shape_vec((1, 1), vec![t_to]).unwrap();
let g = atom.decode_at_coords(t_mat.view()).unwrap();
(0..P_OUT)
.map(|j| scale[j] * g[[0, j]] - z_from[j])
.collect()
};
let unit_nats = 0.5 * categorical_quad_form(&p_from, &dg_raw);
assert!(unit_nats > 0.0, "unit chord must carry Fisher mass");
let target_nats = 0.5 * unit_nats;
let uncertified_metric = metric
.clone()
.with_fisher_factor_kind(FisherFactorKind::UncertifiedApproximation)
.expect("downgrade factor status for refusal test");
let error = steer_to_target_nats(
&term,
&uncertified_metric,
TargetDoseRequest {
atom_k: 0,
metric_row: row,
t_from: &[t_from],
t_to: &[t_to],
target_nats,
config: TargetDoseConfig::default(),
},
None,
)
.expect_err("an uncertified factor cannot solve a full-KL target without a probe");
assert!(matches!(
error,
crate::inference::steering::TargetDoseError::FactorNeedsAppliedDoseProbe {
kind: crate::inference::steering::FisherDoseKind::UncertifiedApproximation
}
));
let seed_plan = steer_to_target_nats(
&term,
&metric,
TargetDoseRequest {
atom_k: 0,
metric_row: row,
t_from: &[t_from],
t_to: &[t_to],
target_nats,
config: TargetDoseConfig::default(),
},
None,
)
.expect("closed-form seed");
assert!(
(seed_plan.steer.predicted_nats.expect("predicted dose") - target_nats).abs()
<= 1e-9 * target_nats,
"closed-form seed dose {} must equal target {target_nats}",
seed_plan.steer.predicted_nats.expect("predicted dose")
);
let expect_a0 = (target_nats / unit_nats).sqrt();
assert!(
(seed_plan.seed_amplitude - expect_a0).abs() <= 1e-9 * expect_a0,
"seed amplitude {} must be sqrt(q*/unit_nats) = {expect_a0}",
seed_plan.seed_amplitude
);
assert!(seed_plan.applied_probe.is_none());
let z_from_probe = z_from.clone();
let p_from_probe = p_from.clone();
let mut probe = move |plan: &SteerPlan| -> Result<AppliedDoseObservation, String> {
let z_to: Vec<f64> = z_from_probe
.iter()
.zip(plan.delta.iter())
.map(|(&z, &d)| z + d)
.collect();
let p_to = softmax(ArrayView1::from(&z_to));
Ok(AppliedDoseObservation {
effective_delta: plan.delta.clone(),
exact_directional_nats: 0.5
* categorical_quad_form(&p_from_probe, plan.delta.as_slice().unwrap()),
measured_nats: kl(&p_from_probe, &p_to),
certified_attainable_upper_nats: None,
})
};
let plan = steer_to_target_nats(
&term,
&metric,
TargetDoseRequest {
atom_k: 0,
metric_row: row,
t_from: &[t_from],
t_to: &[t_to],
target_nats,
config: TargetDoseConfig::default(),
},
Some(&mut probe),
)
.expect("target-dose loop with exact-KL probe");
let measured = plan
.applied_probe
.as_ref()
.expect("applied probe")
.measured_nats;
assert!(
(measured - target_nats).abs() / target_nats <= 2.0e-2,
"measured KL {measured} must land within 2% of target {target_nats}"
);
assert!(
plan.iterations <= 4,
"an in-radius target must converge in a couple of probes; took {}",
plan.iterations
);
assert!(
plan.readout_kl_radius.is_some(),
"an in-radius probe must record a readout-KL radius"
);
}
}