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;
#[derive(Debug, Clone)]
pub struct AmortizedCode {
pub logits: Array2<f64>,
pub coords: Vec<Array2<f64>>,
pub amplitudes: Array2<f64>,
}
#[derive(Debug, Clone)]
pub struct AmortizationErrorStats {
pub coord_rmse: f64,
pub coord_abs_err_quantiles: [f64; 5],
pub gate_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;
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;
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()
&& (log_lambda - last_log_lambda).abs() <= EVIDENCE_REL_TOL * (1.0 + log_lambda.abs())
{
break;
}
last_log_lambda = log_lambda;
}
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>,
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>,
) -> 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()
));
}
let coord_dims: Vec<usize> = coords.iter().map(|c| c.ncols()).collect();
let coord_total: usize = coord_dims.iter().sum();
let t_dim = 2 * k + coord_total;
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 {
for row in 0..n {
targets[[row, offset + axis]] = coord[[row, axis]];
}
}
offset += d;
}
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 (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)?;
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,
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 &d in &self.coord_dims {
let mut block = Array2::<f64>::zeros((m, d));
for axis in 0..d {
for row in 0..m {
block[[row, axis]] = pred[[row, offset + axis]];
}
}
coords.push(block);
offset += d;
}
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>,
) -> 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());
}
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"));
}
for row in 0..pc.nrows() {
for axis in 0..pc.ncols() {
let e = (pc[[row, axis]] - ec[[row, axis]]).abs();
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.logits[[row, atom]] > 0.0;
let e_active = exact_logits[[row, atom]] > 0.0;
if p_active == e_active {
agree += 1;
}
}
}
let gate_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,
gate_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())
.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.gate_agreement > 0.98,
"gate agreement must be near-perfect on a recovered linear map, got {}",
stats.gate_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())
.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)"
);
}
}