use crate::jet_scalar::JetScalar;
use crate::jet_tower::{
KernelChannels, RowProgram, Tower4, program_full_tower, verify_kernel_channels,
};
const LOGB_SIGMA_FLOOR: f64 = 0.01;
#[derive(Clone, Copy, Debug)]
struct GaulssRow {
y: f64,
eta_mu: f64,
eta_ls: f64,
a: f64,
}
struct GaulssLinkRow {
rows: Vec<GaulssRow>,
}
impl RowProgram<2> for GaulssLinkRow {
fn n_rows(&self) -> usize {
self.rows.len()
}
fn primaries(&self, row: usize) -> Result<[f64; 2], String> {
let r = self
.rows
.get(row)
.ok_or_else(|| format!("GaulssLinkRow: row {row} out of range"))?;
Ok([r.eta_mu, r.eta_ls])
}
fn eval<S: JetScalar<2>>(&self, row: usize, p: &[S; 2]) -> Result<S, String> {
let data = self
.rows
.get(row)
.ok_or_else(|| format!("GaulssLinkRow: row {row} out of range"))?;
let eta_mu = &p[0];
let eta_ls = &p[1];
let sigma = eta_ls.exp().add(&S::constant(LOGB_SIGMA_FLOOR));
let r = S::constant(data.y).sub(eta_mu);
let log_term = sigma.ln().scale(data.a);
let r_over_sigma = r.mul(&sigma.recip());
let quad = r_over_sigma.mul(&r_over_sigma).scale(0.5 * data.a);
Ok(log_term.add(&quad))
}
}
fn gaulss_observed_closed_form(row: &GaulssRow) -> KernelChannels<2> {
let s_exp = row.eta_ls.exp();
let sigma = LOGB_SIGMA_FLOOR + s_exp;
let kappa = s_exp / sigma; let kappa_prime = kappa * (1.0 - kappa); let r = row.y - row.eta_mu;
let w = row.a / (sigma * sigma); let m = r * w; let n = r * r * w; let a = row.a;
let value = a * sigma.ln() + 0.5 * a * r * r / (sigma * sigma);
let gradient = [-m, kappa * (a - n)];
let h_mu_mu = w;
let h_mu_ls = 2.0 * m * kappa;
let h_ls_ls = 2.0 * kappa * kappa * n + kappa_prime * (a - n);
let hessian = [[h_mu_mu, h_mu_ls], [h_mu_ls, h_ls_ls]];
KernelChannels {
value,
gradient,
hessian,
third: Vec::new(),
fourth: Vec::new(),
}
}
struct Lcg(u64);
impl Lcg {
fn next_f64(&mut self) -> f64 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((self.0 >> 11) as f64) / ((1u64 << 53) as f64)
}
fn uniform(&mut self, lo: f64, hi: f64) -> f64 {
lo + (hi - lo) * self.next_f64()
}
}
#[test]
fn gaulss_link_jet_tower_matches_production_observed_score_and_hessian() {
let mut rng = Lcg(0x9322_2020_1109_ca75);
let mut rows = Vec::new();
for _ in 0..24 {
rows.push(GaulssRow {
y: rng.uniform(-3.0, 3.0),
eta_mu: rng.uniform(-2.0, 2.0),
eta_ls: rng.uniform(-1.5, 1.5),
a: rng.uniform(0.5, 2.5),
});
}
let program = GaulssLinkRow { rows: rows.clone() };
const REL_TOL: f64 = 1e-11;
for (row, fixture) in rows.iter().enumerate() {
let tower: Box<Tower4<2>> = program_full_tower(&program, row).expect("gaulss jet tower");
let claims = gaulss_observed_closed_form(fixture);
verify_kernel_channels(&tower, &claims, REL_TOL).unwrap_or_else(|e| {
panic!(
"row {row}: gaulss production observed κ-chain-rule tower disagrees with \
#932 jet-tower truth: {e}"
)
});
}
}