use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
use super::{SaeBasisEvaluator, SaeManifoldAtom, SaeManifoldRho, SaeManifoldTerm, Side};
use gam_linalg::faer_ndarray::FaerCholesky;
pub fn decoder_frobenius_norm(decoder: ArrayView2<'_, f64>) -> f64 {
decoder.iter().map(|v| v * v).sum::<f64>().sqrt()
}
pub fn retract_decoder_unit_frobenius(atom: &mut SaeManifoldAtom) -> bool {
let norm = decoder_frobenius_norm(atom.decoder_coefficients.view());
if !(norm.is_finite() && norm > 0.0) {
return false;
}
if (norm - 1.0).abs() <= f64::EPSILON {
return false;
}
atom.absorb_decoder_norm_into_log_amplitude(f64::MIN_POSITIVE);
true
}
pub fn unit_frobenius_tangent_projection(
decoder: ArrayView2<'_, f64>,
ambient_grad: ArrayView2<'_, f64>,
) -> Array2<f64> {
let bb = decoder.iter().map(|v| v * v).sum::<f64>();
let mut out = ambient_grad.to_owned();
if !(bb > 0.0) {
return out;
}
let gb: f64 = ambient_grad
.iter()
.zip(decoder.iter())
.map(|(g, b)| g * b)
.sum();
let coeff = gb / bb;
for (o, b) in out.iter_mut().zip(decoder.iter()) {
*o -= coeff * b;
}
out
}
#[derive(Debug, Clone)]
pub struct LogAmplitudeHoyerEnergy {
pub value: f64,
pub grad: Array1<f64>,
pub hess: Array2<f64>,
}
pub fn log_amplitude_hoyer_energy(s: ArrayView1<'_, f64>, lambda: f64) -> LogAmplitudeHoyerEnergy {
let k = s.len();
let mut grad = Array1::<f64>::zeros(k);
let mut hess = Array2::<f64>::zeros((k, k));
if k <= 1 {
return LogAmplitudeHoyerEnergy {
value: 0.0,
grad,
hess,
};
}
let smax = s.iter().copied().fold(f64::NEG_INFINITY, f64::max);
if !smax.is_finite() {
return LogAmplitudeHoyerEnergy {
value: 0.0,
grad,
hess,
};
}
let a: Vec<f64> = s.iter().map(|&sk| (sk - smax).exp()).collect();
let l1: f64 = a.iter().sum();
let l2_sq: f64 = a.iter().map(|v| v * v).sum();
let l2 = l2_sq.sqrt();
if !(l2 > 0.0 && l1 > 0.0) {
return LogAmplitudeHoyerEnergy {
value: 0.0,
grad,
hess,
};
}
let r = l1 / l2;
let u: Vec<f64> = a.iter().map(|v| v / l2).collect();
let value = lambda * r;
for k1 in 0..k {
grad[k1] = lambda * u[k1] * (1.0 - r * u[k1]);
}
for k1 in 0..k {
for j in 0..k {
let diag = if k1 == j {
u[k1] * (1.0 - 2.0 * r * u[k1])
} else {
0.0
};
let cross = -u[k1] * u[j] * (u[j] + u[k1]) + 3.0 * r * u[k1] * u[k1] * u[j] * u[j];
hess[[k1, j]] = lambda * (diag + cross);
}
}
LogAmplitudeHoyerEnergy { value, grad, hess }
}
pub fn sample_decoded_curve(
evaluator: &dyn SaeBasisEvaluator,
decoder: ArrayView2<'_, f64>,
log_amplitude: f64,
coords: ArrayView1<'_, f64>,
) -> Result<Array2<f64>, String> {
let n = coords.len();
let mut coords2 = Array2::<f64>::zeros((n, 1));
for i in 0..n {
coords2[[i, 0]] = coords[i];
}
let (phi, _jet) = evaluator.evaluate(coords2.view())?;
if phi.ncols() != decoder.nrows() {
return Err(format!(
"sample_decoded_curve: basis width {} != decoder rows {}",
phi.ncols(),
decoder.nrows()
));
}
let mut pts = phi.dot(&decoder);
if log_amplitude != 0.0 {
let amp = log_amplitude.exp();
pts.mapv_inplace(|v| v * amp);
}
Ok(pts)
}
#[derive(Debug, Clone)]
pub struct AffineChartTransition {
pub slope: f64,
pub offset: f64,
pub coord_residual: f64,
pub geometric_residual: f64,
}
impl AffineChartTransition {
pub fn same_manifold(&self, coord_scale: f64, rel_tol: f64) -> bool {
let slope_ok = (self.slope.abs() - 1.0).abs() <= rel_tol;
let coord_ok = coord_scale > 0.0 && self.coord_residual <= rel_tol * coord_scale;
let geom_ok = self.geometric_residual <= rel_tol;
slope_ok && coord_ok && geom_ok
}
}
pub fn affine_chart_transition(
points_a: ArrayView2<'_, f64>,
coords_a: ArrayView1<'_, f64>,
points_b: ArrayView2<'_, f64>,
coords_b: ArrayView1<'_, f64>,
period_a: Option<f64>,
) -> Result<AffineChartTransition, String> {
let (na, p) = points_a.dim();
let (nb, pb) = points_b.dim();
if p != pb {
return Err(format!(
"affine_chart_transition: output dims differ (a: {p}, b: {pb})"
));
}
if na != coords_a.len() || nb != coords_b.len() {
return Err(format!(
"affine_chart_transition: point/coord length mismatch (a: {na} vs {}, b: {nb} vs {})",
coords_a.len(),
coords_b.len()
));
}
if na < 2 || nb < 2 {
return Err("affine_chart_transition: need at least two samples per curve".into());
}
let mut centroid = vec![0.0_f64; p];
for i in 0..na {
for j in 0..p {
centroid[j] += points_a[[i, j]];
}
}
for c in centroid.iter_mut() {
*c /= na as f64;
}
let mut scale_sq = 0.0_f64;
for i in 0..na {
for j in 0..p {
let d = points_a[[i, j]] - centroid[j];
scale_sq += d * d;
}
}
let curve_scale = (scale_sq / na as f64).sqrt();
let mut xs = Vec::with_capacity(nb); let mut ys = Vec::with_capacity(nb); let mut dist_sum = 0.0_f64;
for jb in 0..nb {
let mut best = f64::INFINITY;
let mut best_i = 0usize;
for ia in 0..na {
let mut d = 0.0_f64;
for c in 0..p {
let diff = points_b[[jb, c]] - points_a[[ia, c]];
d += diff * diff;
}
if d < best {
best = d;
best_i = ia;
}
}
dist_sum += best.sqrt();
xs.push(coords_b[jb]);
ys.push(coords_a[best_i]);
}
let geometric_residual = if curve_scale > 0.0 {
(dist_sum / nb as f64) / curve_scale
} else {
f64::INFINITY
};
let mut order: Vec<usize> = (0..nb).collect();
order.sort_by(|&i, &j| {
xs[i]
.partial_cmp(&xs[j])
.unwrap_or(std::cmp::Ordering::Equal)
});
let xo: Vec<f64> = order.iter().map(|&i| xs[i]).collect();
let mut yo: Vec<f64> = order.iter().map(|&i| ys[i]).collect();
if let Some(pp) = period_a {
if pp > 0.0 {
for idx in 1..yo.len() {
let mut d = yo[idx] - yo[idx - 1];
while d > 0.5 * pp {
yo[idx] -= pp;
d -= pp;
}
while d < -0.5 * pp {
yo[idx] += pp;
d += pp;
}
}
}
}
let m = xo.len() as f64;
let mean_x = xo.iter().sum::<f64>() / m;
let mean_y = yo.iter().sum::<f64>() / m;
let mut sxx = 0.0_f64;
let mut sxy = 0.0_f64;
for idx in 0..xo.len() {
let dx = xo[idx] - mean_x;
sxx += dx * dx;
sxy += dx * (yo[idx] - mean_y);
}
if !(sxx > 0.0) {
return Err(
"affine_chart_transition: curve B coordinate has zero spread; slope undefined".into(),
);
}
let slope = sxy / sxx;
let offset = mean_y - slope * mean_x;
let mut resid_sq = 0.0_f64;
for idx in 0..xo.len() {
let pred = slope * xo[idx] + offset;
let e = yo[idx] - pred;
resid_sq += e * e;
}
let coord_residual = (resid_sq / m).sqrt();
Ok(AffineChartTransition {
slope,
offset,
coord_residual,
geometric_residual,
})
}
impl SaeManifoldTerm {
pub fn retract_decoder_gauge_in_loop(&mut self) -> usize {
let mut retracted = 0usize;
for atom in self.atoms.iter_mut() {
if retract_decoder_unit_frobenius(atom) {
retracted += 1;
}
}
retracted
}
pub fn retract_collapsed_decoders_in_loop(&mut self) -> usize {
let k = self.k_atoms();
if k < 2 {
return 0;
}
let mut norms = vec![0.0_f64; k];
for (idx, atom) in self.atoms.iter().enumerate() {
norms[idx] = atom
.decoder_coefficients
.iter()
.map(|v| v * v)
.sum::<f64>()
.sqrt();
}
let mut sorted = norms.clone();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let median = if k % 2 == 1 {
sorted[k / 2]
} else {
0.5 * (sorted[k / 2 - 1] + sorted[k / 2])
};
let reference = norms.iter().copied().fold(0.0_f64, f64::max);
if !(reference > 0.0) {
return 0;
}
let breach_floor = crate::assignment::SAE_ATOM_DECODER_NORM_COLLAPSE_RATIO * reference;
let direction_floor = 1.0e-12 * reference;
let mut retracted = 0usize;
let mut unretractable_zero = 0usize;
for idx in 0..k {
if norms[idx] < breach_floor {
if norms[idx] > direction_floor {
if retract_decoder_unit_frobenius(&mut self.atoms[idx]) {
retracted += 1;
}
} else {
unretractable_zero += 1;
}
}
}
let norms_fmt: Vec<String> = norms.iter().map(|v| format!("{v:.6e}")).collect();
log::warn!(
"[#1939 cone-atom] k={k} norms=[{}] median={median:.6e} reference={reference:.6e} \
breach_floor={breach_floor:.6e} direction_floor={direction_floor:.6e} \
retracted={retracted} unretractable_zero={unretractable_zero}",
norms_fmt.join(", ")
);
retracted
}
pub fn optimize_log_amplitudes_closed_form(
&mut self,
target: ArrayView2<'_, f64>,
rho: &SaeManifoldRho,
) -> Result<(), String> {
let k = self.k_atoms();
if k == 0 {
return Ok(());
}
let saved: Vec<f64> = self.atoms.iter().map(|a| a.log_amplitude).collect();
let df_before = self
.try_fitted_for_rho(rho)
.and_then(|recon| self.data_fit_for_reconstruction(target, recon.view()))
.unwrap_or(f64::INFINITY);
const OFF: f64 = -700.0; let mut designs: Vec<Array2<f64>> = Vec::with_capacity(k);
let mut probe_err: Option<String> = None;
for kk in 0..k {
for (j, atom) in self.atoms.iter_mut().enumerate() {
atom.log_amplitude = if j == kk { 0.0 } else { OFF };
}
match self.try_fitted_for_rho(rho) {
Ok(c) => designs.push(c),
Err(e) => {
probe_err = Some(e);
break;
}
}
}
for (j, atom) in self.atoms.iter_mut().enumerate() {
atom.log_amplitude = saved[j];
}
if let Some(e) = probe_err {
return Err(format!(
"optimize_log_amplitudes_closed_form: per-atom probe fit failed: {e}"
));
}
let metric = self.row_metric();
let whitens = metric.is_some_and(|m| m.whitens_likelihood());
let row_loss_w = self.row_loss_weights();
let n = target.nrows();
let mut gram = Array2::<f64>::zeros((k, k));
let mut rhs = Array1::<f64>::zeros(k);
let mut wdesign: Vec<Vec<f64>> = vec![Vec::new(); k];
for row in 0..n {
let sw = row_loss_w.map_or(1.0, |w| w[row]).sqrt();
let whiten_row = |r: ArrayView1<'_, f64>| -> Vec<f64> {
match metric {
Some(m) if whitens => {
let mut w = m.whiten_residual_row(row, r);
for x in w.iter_mut() {
*x *= sw;
}
w
}
_ => r.iter().map(|&x| x * sw).collect(),
}
};
let wtarget = whiten_row(target.row(row));
for kk in 0..k {
wdesign[kk] = whiten_row(designs[kk].row(row));
}
for j in 0..k {
rhs[j] += wtarget
.iter()
.zip(wdesign[j].iter())
.map(|(t, d)| t * d)
.sum::<f64>();
for kk in j..k {
let g: f64 = wdesign[j]
.iter()
.zip(wdesign[kk].iter())
.map(|(x, y)| x * y)
.sum();
gram[[j, kk]] += g;
if kk != j {
gram[[kk, j]] += g;
}
}
}
}
let scale = (0..k)
.map(|j| gram[[j, j]])
.fold(0.0_f64, f64::max)
.max(1e-300);
for j in 0..k {
gram[[j, j]] += 1e-12 * scale;
}
let raw = gram
.cholesky(Side::Lower)
.map_err(|e| format!("optimize_log_amplitudes_closed_form: cholesky failed: {e}"))?
.solvevec(&rhs);
const AMP_FLOOR: f64 = 1.0e-12;
for atom_idx in 0..k {
let b = raw[atom_idx];
self.atoms[atom_idx].log_amplitude = if b.is_finite() && b > AMP_FLOOR {
b.ln()
} else {
AMP_FLOOR.ln()
};
}
let df_after = self
.try_fitted_for_rho(rho)
.and_then(|recon| self.data_fit_for_reconstruction(target, recon.view()))
.unwrap_or(f64::INFINITY);
let tol = 1.0e-9 * df_before.abs().max(1.0);
if !(df_after <= df_before + tol) {
for (j, atom) in self.atoms.iter_mut().enumerate() {
atom.log_amplitude = saved[j];
}
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AmplitudeConcentration {
SpikeAtSaturation,
Continuous,
Indeterminate,
}
impl AmplitudeConcentration {
pub fn label(self) -> &'static str {
match self {
AmplitudeConcentration::SpikeAtSaturation => "spike_at_saturation",
AmplitudeConcentration::Continuous => "continuous",
AmplitudeConcentration::Indeterminate => "indeterminate",
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct AmplitudeConcentrationCertificate {
pub verdict: AmplitudeConcentration,
pub beta_alpha: f64,
pub beta_beta: f64,
pub log_likelihood: f64,
pub n: usize,
}
impl AmplitudeConcentrationCertificate {
pub fn recommends_radial_axis(&self) -> bool {
matches!(self.verdict, AmplitudeConcentration::Continuous)
}
}
pub fn amplitude_concentration_certificate(
amplitudes: ArrayView1<'_, f64>,
) -> AmplitudeConcentrationCertificate {
let n = amplitudes.len();
let indeterminate = |n: usize| AmplitudeConcentrationCertificate {
verdict: AmplitudeConcentration::Indeterminate,
beta_alpha: f64::NAN,
beta_beta: f64::NAN,
log_likelihood: f64::NAN,
n,
};
if n < 4 {
return indeterminate(n);
}
if amplitudes.iter().any(|a| !a.is_finite() || *a < 0.0) {
return indeterminate(n);
}
let amax = amplitudes.iter().copied().fold(0.0_f64, f64::max);
if !(amax > 0.0) {
return indeterminate(n);
}
let raw: Vec<f64> = amplitudes.iter().map(|&a| (a / amax).clamp(0.0, 1.0)).collect();
let mean_r: f64 = raw.iter().sum::<f64>() / n as f64;
let var_r: f64 = raw.iter().map(|r| (r - mean_r).powi(2)).sum::<f64>() / n as f64;
if !(var_r > f64::EPSILON) {
return indeterminate(n);
}
let nf = n as f64;
let r: Vec<f64> = raw
.iter()
.map(|&x| (x * (nf - 1.0) + 0.5) / nf)
.collect();
let (alpha, beta, loglik) = match fit_beta_mle(&r) {
Some(v) => v,
None => return indeterminate(n),
};
let verdict = if alpha < 1.0 && beta < 1.0 {
AmplitudeConcentration::SpikeAtSaturation
} else {
AmplitudeConcentration::Continuous
};
AmplitudeConcentrationCertificate {
verdict,
beta_alpha: alpha,
beta_beta: beta,
log_likelihood: loglik,
n,
}
}
fn fit_beta_mle(r: &[f64]) -> Option<(f64, f64, f64)> {
let n = r.len();
if n < 2 {
return None;
}
let mut sum_ln = 0.0_f64;
let mut sum_ln1m = 0.0_f64;
let mut mean = 0.0_f64;
let mut mean_sq = 0.0_f64;
for &x in r {
if !(x > 0.0 && x < 1.0) {
return None;
}
sum_ln += x.ln();
sum_ln1m += (1.0 - x).ln();
mean += x;
mean_sq += x * x;
}
let nf = n as f64;
mean /= nf;
let var = (mean_sq / nf - mean * mean).max(f64::EPSILON);
let common = (mean * (1.0 - mean) / var - 1.0).max(1.0e-3);
let mut alpha = (mean * common).max(1.0e-3);
let mut beta = ((1.0 - mean) * common).max(1.0e-3);
let s_ln = sum_ln / nf;
let s_ln1m = sum_ln1m / nf;
for _ in 0..100 {
let psi_ab = digamma(alpha + beta);
let g_a = s_ln - (digamma(alpha) - psi_ab);
let g_b = s_ln1m - (digamma(beta) - psi_ab);
if g_a.abs() < 1.0e-12 && g_b.abs() < 1.0e-12 {
break;
}
let t_ab = trigamma(alpha + beta);
let h_aa = trigamma(alpha) - t_ab;
let h_bb = trigamma(beta) - t_ab;
let h_ab = -t_ab;
let det = h_aa * h_bb - h_ab * h_ab;
if !(det.abs() > 0.0) {
break;
}
let d_a = (h_bb * g_a - h_ab * g_b) / det;
let d_b = (h_aa * g_b - h_ab * g_a) / det;
let mut step = 1.0_f64;
let base = beta_loglik_avg(alpha, beta, s_ln, s_ln1m);
let mut moved = false;
for _ in 0..40 {
let na = alpha + step * d_a;
let nb = beta + step * d_b;
if na > 0.0 && nb > 0.0 && beta_loglik_avg(na, nb, s_ln, s_ln1m) >= base {
alpha = na;
beta = nb;
moved = true;
break;
}
step *= 0.5;
}
if !moved {
break;
}
}
let loglik = nf * beta_loglik_avg(alpha, beta, s_ln, s_ln1m);
if !loglik.is_finite() {
return None;
}
Some((alpha, beta, loglik))
}
fn beta_loglik_avg(alpha: f64, beta: f64, s_ln: f64, s_ln1m: f64) -> f64 {
(alpha - 1.0) * s_ln + (beta - 1.0) * s_ln1m
- (ln_gamma(alpha) + ln_gamma(beta) - ln_gamma(alpha + beta))
}
fn digamma(mut x: f64) -> f64 {
let mut result = 0.0_f64;
while x < 10.0 {
result -= 1.0 / x;
x += 1.0;
}
let inv = 1.0 / x;
let inv2 = inv * inv;
result + x.ln() - 0.5 * inv
- inv2 * (1.0 / 12.0 - inv2 * (1.0 / 120.0 - inv2 / 252.0))
}
fn trigamma(mut x: f64) -> f64 {
let mut result = 0.0_f64;
while x < 10.0 {
result += 1.0 / (x * x);
x += 1.0;
}
let inv = 1.0 / x;
let inv2 = inv * inv;
result
+ inv * (1.0 + inv * (0.5 + inv * (1.0 / 6.0 - inv2 * (1.0 / 30.0 - inv2 / 42.0))))
}
fn ln_gamma(x: f64) -> f64 {
const G: f64 = 7.0;
const C: [f64; 9] = [
0.999_999_999_999_809_93,
676.520_368_121_885_1,
-1_259.139_216_722_402_8,
771.323_428_777_653_13,
-176.615_029_162_140_6,
12.507_343_278_686_905,
-0.138_571_095_265_720_12,
9.984_369_578_019_572e-6,
1.505_632_735_149_311_6e-7,
];
let mut a = C[0];
let t = x + G - 0.5;
for (i, &c) in C.iter().enumerate().skip(1) {
a += c / (x + i as f64 - 1.0);
}
0.5 * (2.0 * std::f64::consts::PI).ln() + (x - 0.5) * t.ln() - t + a.ln()
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::{Array3, array};
fn van_der_corput(n: usize) -> Vec<f64> {
(0..n)
.map(|i| {
let (mut x, mut denom, mut k) = (0.0_f64, 2.0_f64, i + 1);
while k > 0 {
x += (k & 1) as f64 / denom;
denom *= 2.0;
k >>= 1;
}
x
})
.collect()
}
#[test]
fn digamma_trigamma_match_known_values() {
let gamma = 0.577_215_664_901_532_9_f64;
assert!((digamma(1.0) + gamma).abs() < 1.0e-9);
assert!((digamma(2.0) - (1.0 - gamma)).abs() < 1.0e-9);
let pi2_6 = std::f64::consts::PI * std::f64::consts::PI / 6.0;
assert!((trigamma(1.0) - pi2_6).abs() < 1.0e-8);
assert!((ln_gamma(5.0) - 24.0_f64.ln()).abs() < 1.0e-9);
}
#[test]
fn beta_mle_recovers_planted_shape() {
let u = van_der_corput(400);
let samples: Vec<f64> = u.iter().map(|&x| x.sqrt()).collect(); let (a, b, _ll) = fit_beta_mle(&samples).expect("beta fit");
assert!((a - 2.0).abs() < 0.3, "alpha {a}");
assert!((b - 1.0).abs() < 0.3, "beta {b}");
}
#[test]
fn continuous_disk_radius_recommends_radial_axis() {
let u = van_der_corput(500);
let amps = Array1::from_iter(u.iter().map(|&x| x.sqrt()));
let cert = amplitude_concentration_certificate(amps.view());
assert_eq!(cert.verdict, AmplitudeConcentration::Continuous, "{cert:?}");
assert!(cert.recommends_radial_axis());
assert!(cert.beta_alpha > 1.0, "alpha {}", cert.beta_alpha);
}
#[test]
fn true_presence_certifies_spike_at_saturation() {
let jitter = van_der_corput(600);
let amps = Array1::from_iter(jitter.iter().enumerate().map(|(i, &j)| {
let base = if i % 2 == 0 { 0.0 } else { 1.0 };
(base + if base == 0.0 { 0.08 * j } else { -0.08 * j }).clamp(0.0, 1.0)
}));
let cert = amplitude_concentration_certificate(amps.view());
assert_eq!(
cert.verdict,
AmplitudeConcentration::SpikeAtSaturation,
"{cert:?}"
);
assert!(!cert.recommends_radial_axis());
assert!(cert.beta_alpha < 1.0 && cert.beta_beta < 1.0, "{cert:?}");
}
#[test]
fn degenerate_amplitudes_are_indeterminate() {
let flat = Array1::from_elem(50, 0.7);
let cert = amplitude_concentration_certificate(flat.view());
assert_eq!(cert.verdict, AmplitudeConcentration::Indeterminate);
assert!(!cert.recommends_radial_axis());
let zero = Array1::<f64>::zeros(50);
assert_eq!(
amplitude_concentration_certificate(zero.view()).verdict,
AmplitudeConcentration::Indeterminate
);
let few = array![0.1, 0.9];
assert_eq!(
amplitude_concentration_certificate(few.view()).verdict,
AmplitudeConcentration::Indeterminate
);
}
#[derive(Debug)]
struct AffineLineEvaluator;
impl SaeBasisEvaluator for AffineLineEvaluator {
fn evaluate(
&self,
coords: ArrayView2<'_, f64>,
) -> Result<(Array2<f64>, Array3<f64>), String> {
let n = coords.nrows();
let mut phi = Array2::<f64>::zeros((n, 2));
let mut jet = Array3::<f64>::zeros((n, 2, 1));
for i in 0..n {
let t = coords[[i, 0]];
phi[[i, 0]] = 1.0;
phi[[i, 1]] = t;
jet[[i, 0, 0]] = 0.0;
jet[[i, 1, 0]] = 1.0;
}
Ok((phi, jet))
}
fn second_jet_dyn(
&self,
coords: ArrayView2<'_, f64>,
) -> Option<Result<ndarray::Array4<f64>, String>> {
if coords.ncols() != 1 {
return Some(Err(format!(
"AffineLineEvaluator::second_jet_dyn: d = 1 evaluator got {} coords",
coords.ncols()
)));
}
None
}
fn third_jet_dyn(
&self,
coords: ArrayView2<'_, f64>,
) -> Option<Result<ndarray::Array5<f64>, String>> {
if coords.ncols() != 1 {
return Some(Err(format!(
"AffineLineEvaluator::third_jet_dyn: d = 1 evaluator got {} coords",
coords.ncols()
)));
}
None
}
}
#[test]
fn unit_frobenius_tangent_projection_kills_radial_component() {
let b = array![[0.6_f64, 0.0], [0.0, 0.8]]; let radial = b.mapv(|v| 2.5 * v);
let proj = unit_frobenius_tangent_projection(b.view(), radial.view());
let worst = proj.iter().fold(0.0_f64, |a, &v| a.max(v.abs()));
assert!(
worst < 1e-12,
"radial gradient must project to 0, got {worst}"
);
let tangent = array![[0.0_f64, 1.0], [-1.0, 0.0]]; let proj_t = unit_frobenius_tangent_projection(b.view(), tangent.view());
let drift = proj_t
.iter()
.zip(tangent.iter())
.map(|(a, c)| (a - c).abs())
.fold(0.0_f64, f64::max);
assert!(
drift < 1e-12,
"tangent gradient must pass through, drift {drift}"
);
}
#[test]
fn hoyer_energy_gradient_and_hessian_match_fd() {
let s = array![0.3_f64, -0.7, 1.1, 0.05];
let lambda = 1.7_f64;
let base = log_amplitude_hoyer_energy(s.view(), lambda);
let h = 1e-6_f64;
let k = s.len();
for i in 0..k {
let mut sp = s.clone();
sp[i] += h;
let mut sm = s.clone();
sm[i] -= h;
let vp = log_amplitude_hoyer_energy(sp.view(), lambda).value;
let vm = log_amplitude_hoyer_energy(sm.view(), lambda).value;
let fd = (vp - vm) / (2.0 * h);
assert!(
(base.grad[i] - fd).abs() <= 1e-6 * (1.0 + fd.abs()),
"grad[{i}] {} != FD {fd}",
base.grad[i]
);
}
for i in 0..k {
let mut sp = s.clone();
sp[i] += h;
let mut sm = s.clone();
sm[i] -= h;
let gp = log_amplitude_hoyer_energy(sp.view(), lambda).grad;
let gm = log_amplitude_hoyer_energy(sm.view(), lambda).grad;
for j in 0..k {
let fd = (gp[j] - gm[j]) / (2.0 * h);
assert!(
(base.hess[[j, i]] - fd).abs() <= 1e-5 * (1.0 + fd.abs()),
"hess[{j},{i}] {} != FD {fd}",
base.hess[[j, i]]
);
}
}
let shifted = s.mapv(|v| v + 3.4);
let e_shift = log_amplitude_hoyer_energy(shifted.view(), lambda).value;
assert!(
(e_shift - base.value).abs() <= 1e-9 * (1.0 + base.value.abs()),
"Hoyer energy must be invariant to a common amplitude shift"
);
}
#[test]
fn hoyer_energy_prefers_sparse_over_dense() {
let sparse = array![2.0_f64, -3.0, -3.0, -3.0];
let dense = array![0.0_f64, 0.0, 0.0, 0.0];
let es = log_amplitude_hoyer_energy(sparse.view(), 1.0).value;
let ed = log_amplitude_hoyer_energy(dense.view(), 1.0).value;
assert!(es < ed, "sparse energy {es} must be below dense {ed}");
assert!(
(ed - (4.0_f64).sqrt()).abs() < 1e-9,
"dense ratio must be √K"
);
}
#[test]
fn retract_decoder_unit_frobenius_is_image_frozen() {
let coords = array![[0.0_f64], [0.25], [0.5], [0.75], [1.0]];
let ev = AffineLineEvaluator;
let (phi, jet) = ev.evaluate(coords.view()).unwrap();
let decoder = array![[2.0_f64, -1.0], [3.0, 0.5]]; let atom = SaeManifoldAtom::new(
"line",
super::super::SaeAtomBasisKind::Linear,
1,
phi,
jet,
decoder.clone(),
Array2::<f64>::eye(2),
)
.unwrap()
.with_basis_evaluator(std::sync::Arc::new(AffineLineEvaluator));
let before = sample_decoded_curve(
&ev,
atom.decoder_coefficients.view(),
atom.log_amplitude,
coords.column(0),
)
.unwrap();
let mut atom = atom;
let applied = retract_decoder_unit_frobenius(&mut atom);
assert!(applied, "a non-unit decoder must be retracted");
let norm = decoder_frobenius_norm(atom.decoder_coefficients.view());
assert!(
(norm - 1.0).abs() < 1e-12,
"‖B‖_F must be pinned to 1, got {norm}"
);
let after = sample_decoded_curve(
&ev,
atom.decoder_coefficients.view(),
atom.log_amplitude,
coords.column(0),
)
.unwrap();
let drift = before
.iter()
.zip(after.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f64, f64::max);
assert!(
drift < 1e-10,
"retraction must be image-frozen, drift {drift}"
);
assert!(
!retract_decoder_unit_frobenius(&mut atom),
"retraction must be idempotent"
);
}
#[test]
fn affine_transition_detects_same_line_with_reflection_and_offset() {
let ev = AffineLineEvaluator;
let d = array![[0.0_f64, 0.0], [0.6, 0.8]]; let ca = Array1::linspace(0.0, 1.0, 11);
let pts_a = sample_decoded_curve(&ev, d.view(), 0.0, ca.view()).unwrap();
let cb = Array1::linspace(0.0, 1.0, 11);
let db = array![[0.6_f64, 0.8], [-0.6, -0.8]]; let pts_b = sample_decoded_curve(&ev, db.view(), 0.0, cb.view()).unwrap();
let tr = affine_chart_transition(pts_a.view(), ca.view(), pts_b.view(), cb.view(), None)
.unwrap();
assert!(
(tr.slope + 1.0).abs() < 1e-6,
"slope must be -1, got {}",
tr.slope
);
assert!(
(tr.offset - 1.0).abs() < 1e-6,
"offset must be 1, got {}",
tr.offset
);
assert!(
tr.coord_residual < 1e-6,
"coord residual {}",
tr.coord_residual
);
assert!(
tr.geometric_residual < 1e-6,
"geometric residual {}",
tr.geometric_residual
);
assert!(tr.same_manifold(1.0, 1e-3), "must be flagged same-manifold");
}
#[test]
fn affine_transition_rejects_disjoint_curve() {
let ev = AffineLineEvaluator;
let da = array![[0.0_f64, 0.0], [1.0, 0.0]]; let db = array![[0.0_f64, 5.0], [1.0, 0.0]]; let ca = Array1::linspace(0.0, 1.0, 11);
let cb = Array1::linspace(0.0, 1.0, 11);
let pts_a = sample_decoded_curve(&ev, da.view(), 0.0, ca.view()).unwrap();
let pts_b = sample_decoded_curve(&ev, db.view(), 0.0, cb.view()).unwrap();
let tr = affine_chart_transition(pts_a.view(), ca.view(), pts_b.view(), cb.view(), None)
.unwrap();
assert!(
tr.geometric_residual > 1.0,
"disjoint curve must have large geometric residual"
);
assert!(
!tr.same_manifold(1.0, 1e-2),
"disjoint curve must be rejected"
);
}
}