use gam_linalg::faer_ndarray::FaerSvd;
use ndarray::{Array1, Array2, ArrayView2};
const EVIDENCE_REL_TOL: f64 = 1.0e-10;
const EVIDENCE_MAX_ITERS: usize = 500;
const EVIDENCE_REFUSAL_TAIL: usize = 8;
#[derive(Debug, Clone)]
pub struct AmortizedCode {
pub logits: Array2<f64>,
pub coords: Vec<Array2<f64>>,
pub amplitudes: Array2<f64>,
}
pub type AxisPeriods = Vec<Option<f64>>;
#[derive(Debug, Clone)]
pub struct AmortizationErrorStats {
pub coord_rmse: f64,
pub coord_abs_err_quantiles: [f64; 5],
pub support_agreement: f64,
pub amplitude_rmse: f64,
}
pub struct ExactRowSolution<'a> {
pub recon: ArrayView2<'a, f64>,
pub logits: ArrayView2<'a, f64>,
pub coords: &'a [Array2<f64>],
pub amplitudes: ArrayView2<'a, f64>,
}
#[derive(Debug, Clone)]
pub struct AmortizationGap {
pub ev_exact: Option<f64>,
pub ev_amortized: Option<f64>,
pub ev_gap: Option<f64>,
pub errors: AmortizationErrorStats,
pub joint_multistart_fraction: f64,
pub used_quadratic_head: bool,
pub encoder_log_evidence: f64,
pub encoder_feature_dim: usize,
pub encoder_effective_dof: f64,
}
#[derive(Debug, Clone)]
struct Standardizer {
mean: Array1<f64>,
scale: Array1<f64>,
}
impl Standardizer {
fn fit(data: ArrayView2<'_, f64>) -> Self {
let (n, d) = data.dim();
let mut mean = Array1::<f64>::zeros(d);
let mut scale = Array1::<f64>::ones(d);
if n == 0 {
return Self { mean, scale };
}
for col in 0..d {
let mut acc = 0.0;
for row in 0..n {
acc += data[[row, col]];
}
let m = acc / n as f64;
mean[col] = m;
let mut var = 0.0;
for row in 0..n {
let c = data[[row, col]] - m;
var += c * c;
}
let sd = (var / n as f64).sqrt();
scale[col] = if sd > 0.0 && sd.is_finite() { sd } else { 1.0 };
}
Self { mean, scale }
}
fn apply(&self, data: ArrayView2<'_, f64>) -> Array2<f64> {
let (n, d) = data.dim();
let mut out = Array2::<f64>::zeros((n, d));
for row in 0..n {
for col in 0..d {
out[[row, col]] = (data[[row, col]] - self.mean[col]) / self.scale[col];
}
}
out
}
}
#[derive(Debug, Clone)]
enum FeatureMap {
Linear { std: Standardizer },
Quadratic {
raw_std: Standardizer,
feat_std: Standardizer,
},
}
impl FeatureMap {
fn design(&self, x: ArrayView2<'_, f64>) -> Array2<f64> {
match self {
FeatureMap::Linear { std } => std.apply(x),
FeatureMap::Quadratic { raw_std, feat_std } => {
let z = raw_std.apply(x);
let (n, p) = z.dim();
let mut raw = Array2::<f64>::zeros((n, 2 * p));
for row in 0..n {
for col in 0..p {
let v = z[[row, col]];
raw[[row, col]] = v;
raw[[row, p + col]] = v * v;
}
}
feat_std.apply(raw.view())
}
}
}
}
#[derive(Debug, Clone)]
struct EvidenceRidge {
weights: Array2<f64>,
log_evidence: f64,
effective_dof: f64,
}
fn fit_evidence_ridge(
design: ArrayView2<'_, f64>,
targets: ArrayView2<'_, f64>,
) -> Result<EvidenceRidge, String> {
let (n, f_dim) = design.dim();
let t_dim = targets.ncols();
if targets.nrows() != n {
return Err(format!(
"fit_evidence_ridge: design has {n} rows but targets have {}",
targets.nrows()
));
}
if n == 0 || f_dim == 0 || t_dim == 0 {
return Ok(EvidenceRidge {
weights: Array2::zeros((f_dim, t_dim)),
log_evidence: f64::NEG_INFINITY,
effective_dof: 0.0,
});
}
let design_owned = design.to_owned();
let (u_opt, svals, vt_opt) = design_owned
.svd(true, true)
.map_err(|e| format!("fit_evidence_ridge: SVD failed: {e:?}"))?;
let u = u_opt.ok_or_else(|| "fit_evidence_ridge: SVD returned no U".to_string())?;
let vt = vt_opt.ok_or_else(|| "fit_evidence_ridge: SVD returned no Vt".to_string())?;
let r = svals.len();
let z = u.t().dot(&targets); let mut y_energy = vec![0.0_f64; t_dim];
for col in 0..t_dim {
let mut acc = 0.0;
for row in 0..n {
let v = targets[[row, col]];
acc += v * v;
}
y_energy[col] = acc;
}
let mut z_energy = vec![0.0_f64; t_dim]; for col in 0..t_dim {
let mut acc = 0.0;
for i in 0..r {
let v = z[[i, col]];
acc += v * v;
}
z_energy[col] = acc;
}
let s2: Vec<f64> = svals.iter().map(|s| s * s).collect();
let mut alpha = 1.0_f64;
let mut beta = 1.0_f64;
let n_t = (n * t_dim) as f64;
let total_energy: f64 = y_energy.iter().sum();
let energy_floor = (total_energy * f64::EPSILON).max(f64::MIN_POSITIVE);
let mut effective_dof = 0.0_f64;
let mut last_log_lambda = f64::NAN;
let mut evidence_converged = false;
let mut last_delta = f64::NAN;
let mut last_gamma = f64::NAN;
let mut last_w_sq_sum = f64::NAN;
let mut last_rss_sum = f64::NAN;
let mut tail: Vec<(f64, f64)> = Vec::new();
for _ in 0..EVIDENCE_MAX_ITERS {
let lambda = (alpha / beta).max(f64::MIN_POSITIVE);
let mut gamma = 0.0_f64;
for i in 0..r {
gamma += s2[i] / (s2[i] + lambda);
}
let mut w_sq_sum = 0.0_f64;
let mut rss_sum = 0.0_f64;
for col in 0..t_dim {
let mut w_sq = 0.0_f64;
let mut rss_in = 0.0_f64;
for i in 0..r {
let denom = s2[i] + lambda;
let coeff = svals[i] / denom; let zi = z[[i, col]];
w_sq += (coeff * zi) * (coeff * zi);
let shrink = lambda / denom; rss_in += (shrink * zi) * (shrink * zi);
}
w_sq_sum += w_sq;
let tail = (y_energy[col] - z_energy[col]).max(0.0);
rss_sum += rss_in + tail;
}
effective_dof = gamma;
last_gamma = gamma;
last_w_sq_sum = w_sq_sum;
last_rss_sum = rss_sum;
if w_sq_sum <= energy_floor {
evidence_converged = true;
break;
}
alpha = (gamma * t_dim as f64) / w_sq_sum.max(energy_floor);
let well_determined = (n_t - gamma * t_dim as f64).max(f64::MIN_POSITIVE);
beta = well_determined / rss_sum.max(energy_floor);
if !(alpha.is_finite() && beta.is_finite()) {
return Err("fit_evidence_ridge: variance components diverged".to_string());
}
let log_lambda = (alpha / beta).ln();
if last_log_lambda.is_finite() {
last_delta = (log_lambda - last_log_lambda).abs();
}
tail.push((log_lambda, last_delta));
if tail.len() > EVIDENCE_REFUSAL_TAIL {
tail.remove(0);
}
if last_log_lambda.is_finite()
&& (log_lambda - last_log_lambda).abs() <= EVIDENCE_REL_TOL * (1.0 + log_lambda.abs())
{
evidence_converged = true;
break;
}
last_log_lambda = log_lambda;
}
if !evidence_converged {
let log_lambda = (alpha / beta).ln();
let threshold = EVIDENCE_REL_TOL * (1.0 + log_lambda.abs());
let tail_text = tail
.iter()
.map(|(value, delta)| format!("({value:.17e}, {delta:.3e})"))
.collect::<Vec<_>>()
.join(" ");
return Err(format!(
"fit_evidence_ridge: MacKay evidence iteration exhausted its \
{EVIDENCE_MAX_ITERS}-iteration safety cap without meeting the relative \
log-λ fixed-point tolerance {EVIDENCE_REL_TOL:.1e} (last log λ = \
{log_lambda:.17e}, previous = {last_log_lambda:.17e}, |Δ| = {last_delta:.6e} \
against threshold {threshold:.6e}, i.e. {:.3e}× too large; fitted state \
γ = {last_gamma:.6e}, ‖w‖² = {last_w_sq_sum:.6e}, RSS = {last_rss_sum:.6e}, \
energy floor = {energy_floor:.6e}); last {} iterates (log λ, |Δ|) = \
[{tail_text}] — a |Δ| still shrinking across them is a budget shortfall, \
a |Δ| pinned flat is a limit cycle; refusing to mint a non-converged \
evidence ridge — the cap never selects the estimator",
last_delta / threshold,
tail.len()
));
}
let lambda = (alpha / beta).max(f64::MIN_POSITIVE);
let mut rotated = Array2::<f64>::zeros((r, t_dim));
for i in 0..r {
let coeff = svals[i] / (s2[i] + lambda);
for col in 0..t_dim {
rotated[[i, col]] = coeff * z[[i, col]];
}
}
let weights = vt.t().dot(&rotated);
let mut log_det_a = 0.0_f64;
for i in 0..r {
log_det_a += (alpha + beta * s2[i]).ln();
}
log_det_a += (f_dim.saturating_sub(r)) as f64 * alpha.ln();
let two_pi = std::f64::consts::TAU;
let mut w_sq_sum = 0.0_f64;
let mut rss_sum = 0.0_f64;
for col in 0..t_dim {
for i in 0..r {
let denom = s2[i] + lambda;
let coeff = svals[i] / denom;
let zi = z[[i, col]];
w_sq_sum += (coeff * zi) * (coeff * zi);
let shrink = lambda / denom;
rss_sum += (shrink * zi) * (shrink * zi);
}
rss_sum += (y_energy[col] - z_energy[col]).max(0.0);
}
let log_evidence = t_dim as f64
* (0.5 * f_dim as f64 * alpha.ln() + 0.5 * n as f64 * beta.ln()
- 0.5 * log_det_a
- 0.5 * n as f64 * two_pi.ln())
- 0.5 * beta * rss_sum
- 0.5 * alpha * w_sq_sum;
Ok(EvidenceRidge {
weights,
log_evidence,
effective_dof,
})
}
#[derive(Debug, Clone)]
pub struct LearnedAmortizedEncoder {
feature_map: FeatureMap,
weights: Array2<f64>,
target_std: Standardizer,
k_atoms: usize,
coord_dims: Vec<usize>,
coord_periods: Vec<AxisPeriods>,
pub log_evidence: f64,
pub feature_dim: usize,
pub effective_dof: f64,
pub used_quadratic_head: bool,
}
impl LearnedAmortizedEncoder {
fn stack_targets(
logits: ArrayView2<'_, f64>,
coords: &[Array2<f64>],
amplitudes: ArrayView2<'_, f64>,
coord_periods: &[AxisPeriods],
) -> Result<(Array2<f64>, Vec<usize>), String> {
let (n, k) = logits.dim();
if amplitudes.dim() != (n, k) {
return Err(format!(
"LearnedAmortizedEncoder: amplitudes {:?} must match logits ({n}, {k})",
amplitudes.dim()
));
}
if coords.len() != k {
return Err(format!(
"LearnedAmortizedEncoder: {} coord blocks but K={k}",
coords.len()
));
}
if coord_periods.len() != k {
return Err(format!(
"LearnedAmortizedEncoder: {} axis-period blocks but K={k}",
coord_periods.len()
));
}
let coord_dims: Vec<usize> = coords.iter().map(|c| c.ncols()).collect();
let mut coord_target_width = 0usize;
for (atom, &d) in coord_dims.iter().enumerate() {
if coord_periods[atom].len() != d {
return Err(format!(
"LearnedAmortizedEncoder: atom {atom} has {} axis periods but latent dim {d}",
coord_periods[atom].len()
));
}
for axis in 0..d {
if let Some(period) = coord_periods[atom][axis] {
if !(period.is_finite() && period > 0.0) {
return Err(format!(
"LearnedAmortizedEncoder: atom {atom} axis {axis} period {period} must be finite and positive"
));
}
coord_target_width += 2;
} else {
coord_target_width += 1;
}
}
}
let t_dim = 2 * k + coord_target_width;
let mut targets = Array2::<f64>::zeros((n, t_dim));
for col in 0..k {
for row in 0..n {
targets[[row, col]] = logits[[row, col]];
}
}
let mut offset = k;
for (atom, coord) in coords.iter().enumerate() {
if coord.nrows() != n {
return Err(format!(
"LearnedAmortizedEncoder: coord block {atom} has {} rows, expected {n}",
coord.nrows()
));
}
let d = coord_dims[atom];
for axis in 0..d {
match coord_periods[atom][axis] {
Some(period) => {
let w = std::f64::consts::TAU / period;
for row in 0..n {
let ang = coord[[row, axis]] * w;
targets[[row, offset]] = ang.cos();
targets[[row, offset + 1]] = ang.sin();
}
offset += 2;
}
None => {
for row in 0..n {
targets[[row, offset]] = coord[[row, axis]];
}
offset += 1;
}
}
}
}
for col in 0..k {
for row in 0..n {
targets[[row, offset + col]] = amplitudes[[row, col]];
}
}
Ok((targets, coord_dims))
}
pub fn fit(
x: ArrayView2<'_, f64>,
logits: ArrayView2<'_, f64>,
coords: &[Array2<f64>],
amplitudes: ArrayView2<'_, f64>,
) -> Result<Self, String> {
let coord_periods: Vec<AxisPeriods> =
coords.iter().map(|c| vec![None; c.ncols()]).collect();
Self::fit_with_axis_periods(x, logits, coords, amplitudes, &coord_periods)
}
pub fn fit_with_axis_periods(
x: ArrayView2<'_, f64>,
logits: ArrayView2<'_, f64>,
coords: &[Array2<f64>],
amplitudes: ArrayView2<'_, f64>,
coord_periods: &[AxisPeriods],
) -> Result<Self, String> {
let (n, _p) = x.dim();
let k_atoms = logits.ncols();
if n == 0 {
return Err("LearnedAmortizedEncoder::fit: empty training corpus".to_string());
}
let (targets, coord_dims) = Self::stack_targets(logits, coords, amplitudes, coord_periods)?;
let target_std = Standardizer::fit(targets.view());
let targets_std = target_std.apply(targets.view());
let lin_std = Standardizer::fit(x);
let linear_map = FeatureMap::Linear { std: lin_std };
let lin_design = linear_map.design(x);
let lin_fit = fit_evidence_ridge(lin_design.view(), targets_std.view())?;
let raw_std = Standardizer::fit(x);
let z = raw_std.apply(x);
let (nn, p) = z.dim();
let mut raw = Array2::<f64>::zeros((nn, 2 * p));
for row in 0..nn {
for col in 0..p {
let v = z[[row, col]];
raw[[row, col]] = v;
raw[[row, p + col]] = v * v;
}
}
let feat_std = Standardizer::fit(raw.view());
let quad_map = FeatureMap::Quadratic { raw_std, feat_std };
let quad_design = quad_map.design(x);
let quad_fit = fit_evidence_ridge(quad_design.view(), targets_std.view())?;
let use_quadratic = quad_fit.log_evidence > lin_fit.log_evidence;
let (feature_map, fit) = if use_quadratic {
(quad_map, quad_fit)
} else {
(linear_map, lin_fit)
};
let feature_dim = fit.weights.nrows();
Ok(Self {
feature_map,
weights: fit.weights,
target_std,
k_atoms,
coord_dims,
coord_periods: coord_periods.to_vec(),
log_evidence: fit.log_evidence,
feature_dim,
effective_dof: fit.effective_dof,
used_quadratic_head: use_quadratic,
})
}
pub fn predict(&self, x: ArrayView2<'_, f64>) -> Result<AmortizedCode, String> {
let design = self.feature_map.design(x);
if design.ncols() != self.weights.nrows() {
return Err(format!(
"LearnedAmortizedEncoder::predict: design width {} != weight rows {}",
design.ncols(),
self.weights.nrows()
));
}
let m = design.nrows();
let pred_std = design.dot(&self.weights); let t_dim = self.target_std.mean.len();
let mut pred = Array2::<f64>::zeros((m, t_dim));
for row in 0..m {
for col in 0..t_dim {
pred[[row, col]] =
pred_std[[row, col]] * self.target_std.scale[col] + self.target_std.mean[col];
}
}
let k = self.k_atoms;
let mut logits = Array2::<f64>::zeros((m, k));
for col in 0..k {
for row in 0..m {
logits[[row, col]] = pred[[row, col]];
}
}
let mut coords = Vec::with_capacity(k);
let mut offset = k;
for (atom, &d) in self.coord_dims.iter().enumerate() {
let mut block = Array2::<f64>::zeros((m, d));
for axis in 0..d {
match self.coord_periods[atom][axis] {
Some(period) => {
let w = std::f64::consts::TAU / period;
for row in 0..m {
let c = pred[[row, offset]];
let s = pred[[row, offset + 1]];
block[[row, axis]] = (s.atan2(c) / w).rem_euclid(period);
}
offset += 2;
}
None => {
for row in 0..m {
block[[row, axis]] = pred[[row, offset]];
}
offset += 1;
}
}
}
coords.push(block);
}
let mut amplitudes = Array2::<f64>::zeros((m, k));
for col in 0..k {
for row in 0..m {
amplitudes[[row, col]] = pred[[row, offset + col]].max(0.0);
}
}
Ok(AmortizedCode {
logits,
coords,
amplitudes,
})
}
pub fn k_atoms(&self) -> usize {
self.k_atoms
}
pub fn error_stats(
predicted: &AmortizedCode,
exact_logits: ArrayView2<'_, f64>,
exact_coords: &[Array2<f64>],
exact_amplitudes: ArrayView2<'_, f64>,
activity_floor: f64,
) -> Result<AmortizationErrorStats, String> {
let coord_periods: Vec<AxisPeriods> = predicted
.coords
.iter()
.map(|c| vec![None; c.ncols()])
.collect();
Self::error_stats_wrapped(
predicted,
exact_logits,
exact_coords,
exact_amplitudes,
&coord_periods,
activity_floor,
)
}
pub fn error_stats_wrapped(
predicted: &AmortizedCode,
exact_logits: ArrayView2<'_, f64>,
exact_coords: &[Array2<f64>],
exact_amplitudes: ArrayView2<'_, f64>,
coord_periods: &[AxisPeriods],
activity_floor: f64,
) -> Result<AmortizationErrorStats, String> {
let (n, k) = predicted.logits.dim();
if exact_logits.dim() != (n, k) || exact_amplitudes.dim() != (n, k) {
return Err("error_stats: exact logits/amplitudes shape mismatch".to_string());
}
if predicted.coords.len() != k || exact_coords.len() != k {
return Err("error_stats: coord block count mismatch".to_string());
}
if coord_periods.len() != k {
return Err("error_stats: axis-period block count mismatch".to_string());
}
if !(activity_floor.is_finite() && activity_floor >= 0.0) {
return Err(format!(
"error_stats: activity_floor must be finite and non-negative; got {activity_floor}"
));
}
let mut abs_errs: Vec<f64> = Vec::new();
let mut coord_sq = 0.0_f64;
let mut coord_cnt = 0usize;
for atom in 0..k {
let pc = &predicted.coords[atom];
let ec = &exact_coords[atom];
if pc.dim() != ec.dim() {
return Err(format!("error_stats: coord block {atom} shape mismatch"));
}
if coord_periods[atom].len() != pc.ncols() {
return Err(format!(
"error_stats: axis-period width mismatch on atom {atom}"
));
}
for row in 0..pc.nrows() {
for axis in 0..pc.ncols() {
let raw = (pc[[row, axis]] - ec[[row, axis]]).abs();
let e = match coord_periods[atom][axis] {
Some(period) => {
let m = raw.rem_euclid(period);
m.min(period - m)
}
None => raw,
};
abs_errs.push(e);
coord_sq += e * e;
coord_cnt += 1;
}
}
}
let coord_rmse = if coord_cnt > 0 {
(coord_sq / coord_cnt as f64).sqrt()
} else {
0.0
};
let coord_abs_err_quantiles = quantiles_5(&mut abs_errs);
let mut agree = 0usize;
let total = n * k;
for row in 0..n {
for atom in 0..k {
let p_active = predicted.amplitudes[[row, atom]] > activity_floor;
let e_active = exact_amplitudes[[row, atom]] > activity_floor;
if p_active == e_active {
agree += 1;
}
}
}
let support_agreement = if total > 0 {
agree as f64 / total as f64
} else {
1.0
};
let mut amp_sq = 0.0_f64;
for row in 0..n {
for atom in 0..k {
let d = predicted.amplitudes[[row, atom]] - exact_amplitudes[[row, atom]];
amp_sq += d * d;
}
}
let amplitude_rmse = if total > 0 {
(amp_sq / total as f64).sqrt()
} else {
0.0
};
Ok(AmortizationErrorStats {
coord_rmse,
coord_abs_err_quantiles,
support_agreement,
amplitude_rmse,
})
}
}
fn quantiles_5(values: &mut [f64]) -> [f64; 5] {
if values.is_empty() {
return [0.0; 5];
}
values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n = values.len();
let at = |q: f64| -> f64 {
let idx = ((q * (n as f64 - 1.0)).round() as usize).min(n - 1);
values[idx]
};
[values[0], at(0.25), at(0.5), at(0.75), values[n - 1]]
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::Array;
struct Lcg(u64);
impl Lcg {
fn next_f64(&mut self) -> f64 {
self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1);
((self.0 >> 11) as f64) / ((1u64 << 53) as f64)
}
fn normal(&mut self) -> f64 {
let u1 = self.next_f64().max(1.0e-12);
let u2 = self.next_f64();
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
}
#[test]
fn recovers_planted_linear_map_and_keeps_null() {
let mut rng = Lcg(12345);
let n = 400usize;
let p = 6usize;
let k = 3usize; let w_logit = Array::from_shape_fn((p, k), |_| rng.normal());
let w_coord = Array::from_shape_fn((p, k), |_| rng.normal());
let w_amp = Array::from_shape_fn((p, k), |_| rng.normal());
let make = |rng: &mut Lcg, n: usize| {
let x = Array::from_shape_fn((n, p), |_| rng.normal());
let logits = x.dot(&w_logit);
let coords_flat = x.dot(&w_coord);
let amp_raw = x.dot(&w_amp);
let coords: Vec<Array2<f64>> = (0..k)
.map(|a| {
let mut c = Array2::<f64>::zeros((n, 1));
for row in 0..n {
c[[row, 0]] = coords_flat[[row, a]];
}
c
})
.collect();
let amplitudes = amp_raw.mapv(|v| 3.0 + 0.3 * v);
(x, logits, coords, amplitudes)
};
let (x_tr, lg_tr, co_tr, am_tr) = make(&mut rng, n);
let enc = LearnedAmortizedEncoder::fit(x_tr.view(), lg_tr.view(), &co_tr, am_tr.view())
.expect("encoder fits");
let (x_te, lg_te, co_te, am_te) = make(&mut rng, 200);
let code = enc.predict(x_te.view()).expect("predict runs");
let stats =
LearnedAmortizedEncoder::error_stats(&code, lg_te.view(), &co_te, am_te.view(), 0.0)
.expect("stats compute");
let coord_scale = {
let mut s = 0.0;
let mut cnt = 0usize;
for c in &co_te {
for v in c.iter() {
s += v * v;
cnt += 1;
}
}
(s / cnt as f64).sqrt()
};
assert!(
stats.coord_rmse < 0.05 * coord_scale,
"linear map must be recovered: coord_rmse={} vs scale={coord_scale}",
stats.coord_rmse
);
assert!(
stats.support_agreement > 0.98,
"support agreement must be near-perfect on a recovered linear map, got {}",
stats.support_agreement
);
assert!(
!enc.used_quadratic_head,
"the evidence must keep the LINEAR null on linear data (recover-the-null law)"
);
}
#[test]
fn admits_quadratic_head_when_evidence_supports_it() {
let mut rng = Lcg(999);
let n = 500usize;
let p = 4usize;
let k = 1usize;
let w = Array::from_shape_fn((p,), |_| rng.normal());
let make = |rng: &mut Lcg, n: usize| {
let x = Array::from_shape_fn((n, p), |_| rng.normal());
let mut coord = Array2::<f64>::zeros((n, 1));
for row in 0..n {
let mut acc = 0.0;
for j in 0..p {
acc += w[j] * x[[row, j]] * x[[row, j]];
}
coord[[row, 0]] = acc;
}
let logits = Array2::<f64>::from_elem((n, k), 1.0);
let amplitudes = Array2::<f64>::from_elem((n, k), 1.0);
(x, logits, vec![coord], amplitudes)
};
let (x_tr, lg_tr, co_tr, am_tr) = make(&mut rng, n);
let enc = LearnedAmortizedEncoder::fit(x_tr.view(), lg_tr.view(), &co_tr, am_tr.view())
.expect("encoder fits");
assert!(
enc.used_quadratic_head,
"the evidence must admit the quadratic head on genuinely quadratic data"
);
let (x_te, lg_te, co_te, am_te) = make(&mut rng, 200);
let code = enc.predict(x_te.view()).expect("predict runs");
let stats =
LearnedAmortizedEncoder::error_stats(&code, lg_te.view(), &co_te, am_te.view(), 0.0)
.expect("stats");
let coord_var = {
let mut mean = 0.0;
for c in &co_te {
for v in c.iter() {
mean += v;
}
}
mean /= 200.0;
let mut s = 0.0;
for c in &co_te {
for v in c.iter() {
s += (v - mean) * (v - mean);
}
}
(s / 200.0).sqrt()
};
assert!(
stats.coord_rmse < 0.35 * coord_var,
"the quadratic head must fit the quadratic coord well: rmse={} vs sd={coord_var}",
stats.coord_rmse
);
}
#[test]
fn shrinks_to_null_on_pure_noise() {
let mut rng = Lcg(7);
let n = 300usize;
let p = 5usize;
let x = Array::from_shape_fn((n, p), |_| rng.normal());
let mut coord = Array2::<f64>::zeros((n, 1));
for row in 0..n {
coord[[row, 0]] = rng.normal(); }
let logits = Array2::<f64>::from_elem((n, 1), 1.0);
let amplitudes = Array2::<f64>::from_elem((n, 1), 1.0);
let enc = LearnedAmortizedEncoder::fit(
x.view(),
logits.view(),
std::slice::from_ref(&coord),
amplitudes.view(),
)
.expect("fit");
let x_te = Array::from_shape_fn((150, p), |_| rng.normal());
let mut coord_te = Array2::<f64>::zeros((150, 1));
for row in 0..150 {
coord_te[[row, 0]] = rng.normal();
}
let code = enc.predict(x_te.view()).expect("predict");
let mut pmean = 0.0;
for row in 0..150 {
pmean += code.coords[0][[row, 0]];
}
pmean /= 150.0;
let mut pvar = 0.0;
for row in 0..150 {
let d = code.coords[0][[row, 0]] - pmean;
pvar += d * d;
}
pvar /= 150.0;
assert!(
pvar < 0.25,
"on pure noise the encoder must shrink to the mean (pred var={pvar} should be «1)"
);
}
fn ranks(v: &[f64]) -> Vec<f64> {
let mut idx: Vec<usize> = (0..v.len()).collect();
idx.sort_by(|&i, &j| v[i].partial_cmp(&v[j]).unwrap_or(std::cmp::Ordering::Equal));
let mut r = vec![0.0; v.len()];
for (rank, &i) in idx.iter().enumerate() {
r[i] = rank as f64;
}
r
}
fn pearson(a: &[f64], b: &[f64]) -> f64 {
let n = a.len() as f64;
let ma = a.iter().sum::<f64>() / n;
let mb = b.iter().sum::<f64>() / n;
let mut cov = 0.0;
let mut va = 0.0;
let mut vb = 0.0;
for i in 0..a.len() {
let da = a[i] - ma;
let db = b[i] - mb;
cov += da * db;
va += da * da;
vb += db * db;
}
cov / (va.sqrt() * vb.sqrt())
}
fn spearman(a: &[f64], b: &[f64]) -> f64 {
pearson(&ranks(a), &ranks(b))
}
#[test]
fn periodic_axis_recovered_by_embedding_while_raw_is_antipode_biased() {
let mut rng = Lcg(20250708);
let p = 4usize;
let make = |rng: &mut Lcg, n: usize| {
let mut x = Array2::<f64>::zeros((n, p));
let mut t = Array2::<f64>::zeros((n, 1));
for row in 0..n {
let ti = rng.next_f64(); t[[row, 0]] = ti;
let ang = std::f64::consts::TAU * ti;
x[[row, 0]] = ang.cos() + 0.01 * rng.normal();
x[[row, 1]] = ang.sin() + 0.01 * rng.normal();
x[[row, 2]] = 0.01 * rng.normal();
x[[row, 3]] = 0.01 * rng.normal();
}
let logits = Array2::<f64>::from_elem((n, 1), 1.0);
let amplitudes = Array2::<f64>::from_elem((n, 1), 1.0);
(x, logits, vec![t], amplitudes)
};
let (x_tr, lg_tr, co_tr, am_tr) = make(&mut rng, 500);
let periods: Vec<AxisPeriods> = vec![vec![Some(1.0)]];
let enc = LearnedAmortizedEncoder::fit_with_axis_periods(
x_tr.view(),
lg_tr.view(),
&co_tr,
am_tr.view(),
&periods,
)
.expect("embedding encoder fits");
let enc_raw = LearnedAmortizedEncoder::fit(x_tr.view(), lg_tr.view(), &co_tr, am_tr.view())
.expect("raw encoder fits");
let (x_te, lg_te, co_te, am_te) = make(&mut rng, 300);
let code = enc.predict(x_te.view()).expect("embedding predict");
let code_raw = enc_raw.predict(x_te.view()).expect("raw predict");
for row in 0..code.coords[0].nrows() {
let th = code.coords[0][[row, 0]];
assert!(
(0.0..1.0).contains(&th),
"embedding-recovered coord must be wrapped into [0,1), got {th}"
);
}
let stats = LearnedAmortizedEncoder::error_stats_wrapped(
&code,
lg_te.view(),
&co_te,
am_te.view(),
&periods,
0.0,
)
.expect("wrapped stats");
let stats_raw = LearnedAmortizedEncoder::error_stats_wrapped(
&code_raw,
lg_te.view(),
&co_te,
am_te.view(),
&periods,
0.0,
)
.expect("wrapped stats raw");
assert!(
stats.coord_rmse < 0.05,
"embedding regression must recover the circle: wrapped rmse={} (< 0.05)",
stats.coord_rmse
);
assert!(
stats_raw.coord_rmse > 0.15,
"raw regression must be antipode-biased on the in-cloud seam: wrapped rmse={} (> 0.15)",
stats_raw.coord_rmse
);
assert!(
stats.coord_rmse < 0.25 * stats_raw.coord_rmse,
"embedding must beat raw by a wide margin: {} vs {}",
stats.coord_rmse,
stats_raw.coord_rmse
);
let t_true: Vec<f64> = (0..co_te[0].nrows()).map(|r| co_te[0][[r, 0]]).collect();
let t_hat: Vec<f64> = (0..code.coords[0].nrows())
.map(|r| code.coords[0][[r, 0]])
.collect();
let rho = spearman(&t_true, &t_hat);
assert!(
rho > 0.9,
"embedding-recovered coordinate must be rank-consistent with truth: spearman={rho}"
);
}
#[test]
fn non_periodic_axis_is_bit_identical_raw_passthrough() {
let mut rng = Lcg(4242);
let n = 300usize;
let p = 5usize;
let k = 2usize;
let w_coord = Array::from_shape_fn((p, k), |_| 3.0 * rng.normal());
let make = |rng: &mut Lcg, n: usize| {
let x = Array::from_shape_fn((n, p), |_| rng.normal());
let coords_flat = x.dot(&w_coord);
let coords: Vec<Array2<f64>> = (0..k)
.map(|a| {
let mut c = Array2::<f64>::zeros((n, 1));
for row in 0..n {
c[[row, 0]] = coords_flat[[row, a]];
}
c
})
.collect();
let logits = Array2::<f64>::from_elem((n, k), 1.0);
let amplitudes = Array2::<f64>::from_elem((n, k), 1.0);
(x, logits, coords, amplitudes)
};
let (x_tr, lg_tr, co_tr, am_tr) = make(&mut rng, n);
let enc_default =
LearnedAmortizedEncoder::fit(x_tr.view(), lg_tr.view(), &co_tr, am_tr.view())
.expect("default fit");
let all_none: Vec<AxisPeriods> = co_tr.iter().map(|c| vec![None; c.ncols()]).collect();
let enc_none = LearnedAmortizedEncoder::fit_with_axis_periods(
x_tr.view(),
lg_tr.view(),
&co_tr,
am_tr.view(),
&all_none,
)
.expect("all-none fit");
let (x_te, ..) = make(&mut rng, 120);
let code_default = enc_default.predict(x_te.view()).expect("predict default");
let code_none = enc_none.predict(x_te.view()).expect("predict none");
let mut saw_out_of_unit = false;
for atom in 0..k {
for row in 0..code_default.coords[atom].nrows() {
let a = code_default.coords[atom][[row, 0]];
let b = code_none.coords[atom][[row, 0]];
assert!(
a == b,
"all-None periods must reproduce the raw fit bit-for-bit: {a} vs {b}"
);
if a.abs() > 1.0 {
saw_out_of_unit = true;
}
}
}
assert!(
saw_out_of_unit,
"a flat axis must pass through raw (out-of-[0,1) coords survive, not atan2-wrapped)"
);
}
#[test]
fn wrap_aware_scoring_charges_the_short_arc() {
let predicted = AmortizedCode {
logits: Array2::<f64>::from_elem((1, 1), 1.0),
coords: vec![Array2::from_shape_vec((1, 1), vec![0.98]).unwrap()],
amplitudes: Array2::<f64>::from_elem((1, 1), 1.0),
};
let exact_logits = Array2::<f64>::from_elem((1, 1), 1.0);
let exact_coords = vec![Array2::from_shape_vec((1, 1), vec![0.02]).unwrap()];
let exact_amp = Array2::<f64>::from_elem((1, 1), 1.0);
let periods: Vec<AxisPeriods> = vec![vec![Some(1.0)]];
let wrapped = LearnedAmortizedEncoder::error_stats_wrapped(
&predicted,
exact_logits.view(),
&exact_coords,
exact_amp.view(),
&periods,
0.0,
)
.expect("wrapped stats");
assert!(
(wrapped.coord_rmse - 0.04).abs() < 1.0e-12,
"0.98 vs 0.02 on period 1 must score as the 0.04 short arc, got {}",
wrapped.coord_rmse
);
assert!(
(wrapped.coord_abs_err_quantiles[4] - 0.04).abs() < 1.0e-12,
"the max wrapped abs error must be the 0.04 short arc, got {}",
wrapped.coord_abs_err_quantiles[4]
);
let raw = LearnedAmortizedEncoder::error_stats(
&predicted,
exact_logits.view(),
&exact_coords,
exact_amp.view(),
0.0,
)
.expect("raw stats");
assert!(
(raw.coord_rmse - 0.96).abs() < 1.0e-12,
"the unwrapped scorer must charge the full 0.96, got {}",
raw.coord_rmse
);
}
}