use gam_problem::RowMetric;
use ndarray::{Array1, Array2};
use std::sync::Arc;
fn softmax(z: &[f64]) -> Vec<f64> {
let max_z = z.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let exps: Vec<f64> = z.iter().map(|&v| (v - max_z).exp()).collect();
let sum: f64 = exps.iter().sum();
exps.iter().map(|&v| v / sum).collect()
}
fn kl(p: &[f64], q: &[f64]) -> f64 {
p.iter()
.zip(q.iter())
.map(|(&pi, &qi)| if pi > 0.0 { pi * (pi / qi).ln() } else { 0.0 })
.sum()
}
fn categorical_fisher_row_metric(p_probs: &[f64]) -> RowMetric {
let k = p_probs.len();
let mut flat = vec![0.0_f64; k * k];
for c in 0..k {
let sqrt_pc = p_probs[c].sqrt();
for i in 0..k {
let e_ci = if i == c { 1.0 } else { 0.0 };
flat[i * k + c] = sqrt_pc * (e_ci - p_probs[i]);
}
}
let u = Array2::from_shape_vec((1, k * k), flat)
.expect("flat was built with exactly k*k entries just above");
RowMetric::output_fisher(Arc::new(u), k, k)
.expect("the (1, k*k) row factor matches the declared p = rank = k")
.with_fisher_factor_kind(gam_problem::FisherFactorKind::ExactFull)
.expect("a square k x k factor is an admissible exact-full Fisher factor")
}
fn analytic_fisher_quad_form(p_probs: &[f64], delta: &[f64]) -> f64 {
let k = p_probs.len();
let mut acc = 0.0;
for i in 0..k {
for j in 0..k {
let f_ij = if i == j {
p_probs[i] - p_probs[i] * p_probs[j]
} else {
-p_probs[i] * p_probs[j]
};
acc += delta[i] * f_ij * delta[j];
}
}
acc
}
#[test]
fn fisher_mass_units_match_analytic_quadratic_form() {
let p_probs = [0.5_f64, 0.3, 0.2];
let delta = [0.01_f64, -0.02, 0.005];
let metric = categorical_fisher_row_metric(&p_probs);
assert_eq!(metric.p_out(), 3);
assert_eq!(metric.n_rows(), 1);
let delta_arr = Array1::from_vec(delta.to_vec());
let predicted_nats = 0.5 * metric.fisher_mass(0, delta_arr.view());
let analytic_nats = 0.5 * analytic_fisher_quad_form(&p_probs, &delta);
assert!(
analytic_nats > 0.0,
"sanity: quadratic form must be positive"
);
let rel_err = (predicted_nats - analytic_nats).abs() / analytic_nats.abs();
assert!(
rel_err < 1e-10,
"fisher_mass-based predicted_nats {predicted_nats} vs analytic 0.5*deltaT*F*delta \
{analytic_nats}: rel_err {rel_err} >= 1e-10"
);
}
#[test]
fn fisher_mass_quadratic_matches_exact_kl_at_small_delta() {
let p_probs = [0.5_f64, 0.3, 0.2];
let z: Vec<f64> = p_probs.iter().map(|p| p.ln()).collect();
let dir = [1.0_f64, -1.7, 0.4];
let metric = categorical_fisher_row_metric(&p_probs);
let mut last_rel_err = f64::INFINITY;
for &eps in &[1e-2_f64, 1e-3, 1e-4] {
let delta: Vec<f64> = dir.iter().map(|d| eps * d).collect();
let z_perturbed: Vec<f64> = z.iter().zip(delta.iter()).map(|(zi, di)| zi + di).collect();
let q_probs = softmax(&z_perturbed);
let exact_kl = kl(&p_probs, &q_probs);
let delta_arr = Array1::from_vec(delta.clone());
let predicted_nats = 0.5 * metric.fisher_mass(0, delta_arr.view());
assert!(
exact_kl > 0.0,
"sanity: exact KL must be positive for nonzero delta"
);
last_rel_err = (predicted_nats - exact_kl).abs() / exact_kl.abs();
}
assert!(
last_rel_err < 1e-2,
"quadratic fisher_mass prediction vs exact KL at small delta: \
rel_err {last_rel_err} >= 1% — the fisher_mass units are not KL-nats"
);
}