use crate::jet_scalar::JetScalar;
use crate::jet_tower::{
KernelChannels, RowProgram, Tower4, program_full_tower, verify_kernel_channels,
};
#[derive(Clone, Copy, Debug)]
struct CauseRow {
eta1: f64,
eta0: f64,
s: f64,
w: f64,
delta: f64,
has_entry: bool,
}
struct CauseSpecificRow {
rows: Vec<CauseRow>,
}
impl RowProgram<3> for CauseSpecificRow {
fn n_rows(&self) -> usize {
self.rows.len()
}
fn primaries(&self, row: usize) -> Result<[f64; 3], String> {
let r = self
.rows
.get(row)
.ok_or_else(|| format!("CauseSpecificRow: row {row} out of range"))?;
Ok([r.eta1, r.eta0, r.s])
}
fn eval<S: JetScalar<3>>(&self, row: usize, p: &[S; 3]) -> Result<S, String> {
let data = self
.rows
.get(row)
.ok_or_else(|| format!("CauseSpecificRow: row {row} out of range"))?;
let eta1 = &p[0];
let eta0 = &p[1];
let s = &p[2];
let mut ell = eta1.exp();
if data.has_entry {
ell = ell.sub(&eta0.exp());
}
if data.delta != 0.0 {
ell = ell.sub(&eta1.add(&s.ln()).scale(data.delta));
}
Ok(ell.scale(data.w))
}
}
fn cause_specific_closed_form(
row: &CauseRow,
third_dirs: &[[f64; 3]],
fourth_pairs: &[([f64; 3], [f64; 3])],
) -> KernelChannels<3> {
let w = row.w;
let e1 = row.eta1.exp();
let e0 = row.eta0.exp();
let entry = if row.has_entry { 1.0 } else { 0.0 };
let s = row.s;
let d = row.delta;
let value = w * (e1 - entry * e0 - d * (row.eta1 + s.ln()));
let gradient = [w * (e1 - d), -w * entry * e0, -w * d / s];
let h_diag = [w * e1, -w * entry * e0, w * d / (s * s)];
let t3_diag = [w * e1, -w * entry * e0, -2.0 * w * d / (s * s * s)];
let t4_diag = [w * e1, -w * entry * e0, 6.0 * w * d / (s * s * s * s)];
let mut hessian = [[0.0_f64; 3]; 3];
for a in 0..3 {
hessian[a][a] = h_diag[a];
}
let third = third_dirs
.iter()
.map(|dir| {
let mut m = [[0.0_f64; 3]; 3];
for a in 0..3 {
m[a][a] = t3_diag[a] * dir[a];
}
(*dir, m)
})
.collect();
let fourth = fourth_pairs
.iter()
.map(|(u, v)| {
let mut m = [[0.0_f64; 3]; 3];
for a in 0..3 {
m[a][a] = t4_diag[a] * u[a] * v[a];
}
(*u, *v, m)
})
.collect();
KernelChannels {
value,
gradient,
hessian,
third,
fourth,
}
}
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 cause_specific_royston_parmar_jet_tower_matches_production_directional_weights() {
let mut rng = Lcg(0x9322_2020_1109_5171);
let third_dirs: [[f64; 3]; 3] = [[0.7, -1.3, 0.5], [-0.4, 0.6, -0.9], [1.2, 0.2, 0.3]];
let fourth_pairs: [([f64; 3], [f64; 3]); 3] = [
([0.7, -1.3, 0.5], [-0.4, 0.6, -0.9]),
([-0.4, 0.6, -0.9], [1.2, 0.2, 0.3]),
([1.2, 0.2, 0.3], [0.7, -1.3, 0.5]),
];
let mut rows = Vec::new();
for i in 0..24 {
rows.push(CauseRow {
eta1: rng.uniform(-1.5, 1.5),
eta0: rng.uniform(-1.5, 1.5),
s: rng.uniform(0.2, 3.0),
w: rng.uniform(0.4, 2.5),
delta: if i % 2 == 0 { 1.0 } else { 0.0 },
has_entry: i % 3 != 0,
});
}
let program = CauseSpecificRow { rows: rows.clone() };
const REL_TOL: f64 = 1e-11;
for (row, fixture) in rows.iter().enumerate() {
let tower: Box<Tower4<3>> =
program_full_tower(&program, row).expect("cause-specific jet tower");
let claims = cause_specific_closed_form(fixture, &third_dirs, &fourth_pairs);
verify_kernel_channels(&tower, &claims, REL_TOL).unwrap_or_else(|e| {
panic!(
"row {row}: cause-specific Royston-Parmar production directional-derivative \
weights disagree with #932 jet-tower truth: {e}"
)
});
}
}