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
}
}
#[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 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)"
);
}
#[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)"
);
}
}