use ndarray::{Array1, Array2, Array3, ArrayView2, ArrayView3};
const RESIDUAL_EM_ENERGY_FLOOR: f64 = 1e-12;
#[must_use]
pub fn residual_em_score(
x: ArrayView2<'_, f64>,
per_atom_recon: ArrayView3<'_, f64>,
nonneg: bool,
) -> (Array2<f64>, Array2<f64>) {
let (n, f, d) = per_atom_recon.dim();
assert_eq!(
x.dim(),
(n, d),
"residual_em_score: x is {:?} but per_atom_recon is {:?}",
x.dim(),
(n, f, d)
);
let mut code = Array2::<f64>::zeros((n, f));
let mut relative_residual = Array2::<f64>::zeros((n, f));
for row in 0..n {
let mut row_energy = 0.0;
for j in 0..d {
let xj = x[[row, j]];
row_energy += xj * xj;
}
let row_scale = row_energy.max(RESIDUAL_EM_ENERGY_FLOOR);
for atom in 0..f {
let mut rr = 0.0;
let mut rx = 0.0;
for j in 0..d {
let rj = per_atom_recon[[row, atom, j]];
rr += rj * rj;
rx += rj * x[[row, j]];
}
let denom = rr.max(RESIDUAL_EM_ENERGY_FLOOR);
let s = rx / denom;
let c = if nonneg { s.max(0.0) } else { s };
let mut resid = 0.0;
for j in 0..d {
let e = c * per_atom_recon[[row, atom, j]] - x[[row, j]];
resid += e * e;
}
code[[row, atom]] = c;
relative_residual[[row, atom]] = resid / row_scale;
}
}
(code, relative_residual)
}
#[must_use]
pub fn residual_em_score_vjp(
x: ArrayView2<'_, f64>,
per_atom_recon: ArrayView3<'_, f64>,
nonneg: bool,
g_code: ArrayView2<'_, f64>,
g_relative_residual: ArrayView2<'_, f64>,
) -> Array3<f64> {
let (n, f, d) = per_atom_recon.dim();
assert_eq!(x.dim(), (n, d), "residual_em_score_vjp: x shape mismatch");
assert_eq!(g_code.dim(), (n, f), "residual_em_score_vjp: g_code shape");
assert_eq!(
g_relative_residual.dim(),
(n, f),
"residual_em_score_vjp: g_relative_residual shape"
);
let mut grad = Array3::<f64>::zeros((n, f, d));
for row in 0..n {
let mut row_energy = 0.0;
for j in 0..d {
let xj = x[[row, j]];
row_energy += xj * xj;
}
let row_scale = row_energy.max(RESIDUAL_EM_ENERGY_FLOOR);
for atom in 0..f {
let mut rr = 0.0;
let mut rx = 0.0;
for j in 0..d {
let rj = per_atom_recon[[row, atom, j]];
rr += rj * rj;
rx += rj * x[[row, j]];
}
let denom_active = rr >= RESIDUAL_EM_ENERGY_FLOOR;
let denom = rr.max(RESIDUAL_EM_ENERGY_FLOOR);
let s = rx / denom;
let active = if nonneg { s >= 0.0 } else { true };
let c = if nonneg { s.max(0.0) } else { s };
let mut e_dot_r = 0.0;
for j in 0..d {
let e = c * per_atom_recon[[row, atom, j]] - x[[row, j]];
e_dot_r += e * per_atom_recon[[row, atom, j]];
}
let gc = g_code[[row, atom]];
let gq = g_relative_residual[[row, atom]];
let a = if active {
gc + 2.0 * gq * e_dot_r / row_scale
} else {
0.0
};
let coeff_e = 2.0 * gq * c / row_scale;
for j in 0..d {
let rj = per_atom_recon[[row, atom, j]];
let ds_drj = if denom_active {
(x[[row, j]] - 2.0 * s * rj) / denom
} else {
x[[row, j]] / denom
};
let e = c * rj - x[[row, j]];
grad[[row, atom, j]] = a * ds_drj + coeff_e * e;
}
}
}
grad
}
#[derive(Debug, Clone)]
pub enum SaeCriterionAtom {
DataFitPriors {
value: f64,
grad: Array1<f64>,
},
LaplaceComplexity {
value: f64,
grad: Array1<f64>,
},
Occam {
value: f64,
grad: Array1<f64>,
},
ImplicitStationarityCorrection {
grad: Array1<f64>,
},
}
impl SaeCriterionAtom {
#[must_use]
pub fn value(&self) -> f64 {
match self {
Self::DataFitPriors { value, .. }
| Self::LaplaceComplexity { value, .. }
| Self::Occam { value, .. } => *value,
Self::ImplicitStationarityCorrection { .. } => 0.0,
}
}
#[must_use]
pub fn grad(&self) -> &Array1<f64> {
match self {
Self::DataFitPriors { grad, .. }
| Self::LaplaceComplexity { grad, .. }
| Self::Occam { grad, .. }
| Self::ImplicitStationarityCorrection { grad } => grad,
}
}
#[must_use]
pub fn label(&self) -> &'static str {
match self {
Self::DataFitPriors { .. } => "data_fit_priors",
Self::LaplaceComplexity { .. } => "laplace_complexity",
Self::Occam { .. } => "occam",
Self::ImplicitStationarityCorrection { .. } => "implicit_stationarity_correction",
}
}
}
#[derive(Debug, Clone)]
pub struct SaeCriterion {
atoms: Vec<SaeCriterionAtom>,
n_rho: usize,
}
impl SaeCriterion {
#[must_use]
pub fn assemble(
data_fit_priors_value: f64,
laplace_complexity_value: f64,
occam: f64,
explicit: Array1<f64>,
logdet_trace: Array1<f64>,
occam_grad: Array1<f64>,
implicit_correction: Array1<f64>,
) -> Self {
let n_rho = explicit.len();
let atoms = vec![
SaeCriterionAtom::DataFitPriors {
value: data_fit_priors_value,
grad: explicit,
},
SaeCriterionAtom::LaplaceComplexity {
value: laplace_complexity_value,
grad: logdet_trace,
},
SaeCriterionAtom::Occam {
value: -occam,
grad: occam_grad,
},
SaeCriterionAtom::ImplicitStationarityCorrection {
grad: implicit_correction,
},
];
Self { atoms, n_rho }
}
#[must_use]
pub fn value(&self) -> f64 {
self.atoms.iter().map(SaeCriterionAtom::value).sum()
}
#[must_use]
pub fn gradient(&self) -> Array1<f64> {
let mut out = Array1::<f64>::zeros(self.n_rho);
for atom in &self.atoms {
out += atom.grad();
}
out
}
#[must_use]
pub fn atoms(&self) -> &[SaeCriterionAtom] {
&self.atoms
}
#[must_use]
pub fn n_rho(&self) -> usize {
self.n_rho
}
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
fn sample_criterion() -> SaeCriterion {
SaeCriterion::assemble(
3.0, 1.0, 0.5, array![0.10, -0.20, 0.05], array![0.01, 0.02, -0.03], array![-0.04, 0.00, 0.06], array![0.07, -0.01, 0.00], )
}
#[test]
fn value_is_atom_sum() {
let crit = sample_criterion();
let expected = 3.0 + 1.0 - 0.5;
assert!((crit.value() - expected).abs() < 1e-12);
let by_atom: f64 = crit.atoms().iter().map(SaeCriterionAtom::value).sum();
assert!((by_atom - expected).abs() < 1e-12);
}
#[test]
fn gradient_is_channel_sum_including_correction() {
let crit = sample_criterion();
let g = crit.gradient();
let expected = array![
0.10 + 0.01 - 0.04 + 0.07,
-0.20 + 0.02 + 0.00 - 0.01,
0.05 - 0.03 + 0.06 + 0.00
];
for i in 0..3 {
assert!(
(g[i] - expected[i]).abs() < 1e-12,
"coord {i}: {} vs {}",
g[i],
expected[i]
);
}
}
#[test]
fn implicit_correction_atom_is_gradient_only() {
let atom = SaeCriterionAtom::ImplicitStationarityCorrection {
grad: array![1.0, 2.0, 3.0],
};
assert_eq!(atom.value(), 0.0);
assert_eq!(atom.grad().sum(), 6.0);
assert_eq!(atom.label(), "implicit_stationarity_correction");
}
#[test]
fn residual_em_score_matches_hand_computation() {
let x = array![[2.0, 1.0]];
let recon = ndarray::Array3::from_shape_vec((1, 1, 2), vec![1.0, 0.0]).unwrap();
for nonneg in [true, false] {
let (code, relres) = residual_em_score(x.view(), recon.view(), nonneg);
assert!((code[[0, 0]] - 2.0).abs() < 1e-12, "code ({nonneg})");
assert!((relres[[0, 0]] - 0.2).abs() < 1e-12, "relres ({nonneg})");
let g_c = array![[1.0]];
let g_q = array![[0.0]];
let grad =
residual_em_score_vjp(x.view(), recon.view(), nonneg, g_c.view(), g_q.view());
assert!(
(grad[[0, 0, 0]] - (-2.0)).abs() < 1e-12,
"dc/dr0 ({nonneg})"
);
assert!((grad[[0, 0, 1]] - 1.0).abs() < 1e-12, "dc/dr1 ({nonneg})");
let g_c = array![[0.0]];
let g_q = array![[1.0]];
let grad =
residual_em_score_vjp(x.view(), recon.view(), nonneg, g_c.view(), g_q.view());
assert!((grad[[0, 0, 0]] - 0.0).abs() < 1e-12, "dq/dr0 ({nonneg})");
assert!(
(grad[[0, 0, 1]] - (-0.8)).abs() < 1e-12,
"dq/dr1 ({nonneg})"
);
}
}
#[test]
fn residual_em_score_clamp_kills_gradient_but_signed_survives() {
let x = array![[2.0, 1.0]];
let recon = ndarray::Array3::from_shape_vec((1, 1, 2), vec![-1.0, 0.0]).unwrap();
let g_c = array![[1.0]];
let g_q = array![[1.0]];
let (code_nn, _) = residual_em_score(x.view(), recon.view(), true);
assert!((code_nn[[0, 0]] - 0.0).abs() < 1e-12, "clamped code is 0");
let grad_nn = residual_em_score_vjp(x.view(), recon.view(), true, g_c.view(), g_q.view());
assert!(grad_nn[[0, 0, 0]].abs() < 1e-12, "clamped grad r0 = 0");
assert!(grad_nn[[0, 0, 1]].abs() < 1e-12, "clamped grad r1 = 0");
let (code_sg, _) = residual_em_score(x.view(), recon.view(), false);
assert!((code_sg[[0, 0]] - (-2.0)).abs() < 1e-12, "signed code = -2");
let grad_sg = residual_em_score_vjp(x.view(), recon.view(), false, g_c.view(), g_q.view());
assert!(
grad_sg[[0, 0, 0]].abs() + grad_sg[[0, 0, 1]].abs() > 1e-9,
"signed grad is live"
);
}
#[test]
fn atoms_have_distinct_labels() {
let crit = sample_criterion();
let labels: Vec<&str> = crit.atoms().iter().map(SaeCriterionAtom::label).collect();
let mut sorted = labels.clone();
sorted.sort_unstable();
sorted.dedup();
assert_eq!(sorted.len(), labels.len(), "labels must be distinct");
}
}