#![cfg(test)]
#![cfg(test)]
use crate::jet_scalar::JetScalar;
use crate::jet_tower::{
KernelChannels, RowProgram, Tower4, program_full_tower, verify_kernel_channels,
};
#[derive(Clone, Copy, Debug)]
struct MultRow<const M: usize> {
eta: [f64; M],
obs: usize,
w: f64,
}
struct MultinomialSoftmaxRow<const M: usize> {
rows: Vec<MultRow<M>>,
}
impl<const M: usize> RowProgram<M> for MultinomialSoftmaxRow<M> {
fn n_rows(&self) -> usize {
self.rows.len()
}
fn primaries(&self, row: usize) -> Result<[f64; M], String> {
let r = self
.rows
.get(row)
.ok_or_else(|| format!("MultinomialSoftmaxRow: row {row} out of range"))?;
Ok(r.eta)
}
fn eval<S: JetScalar<M>>(&self, row: usize, p: &[S; M]) -> Result<S, String> {
let data = self
.rows
.get(row)
.ok_or_else(|| format!("MultinomialSoftmaxRow: row {row} out of range"))?;
let mut z = S::constant(1.0);
for a in 0..M {
z = z.add(&p[a].exp());
}
let mut ell = z.ln().scale(data.w);
if data.obs < M {
ell = ell.sub(&p[data.obs].scale(data.w));
}
Ok(ell)
}
}
fn softmax_active<const M: usize>(eta: &[f64; M]) -> [f64; M] {
let mut z = 1.0_f64;
let mut ex = [0.0_f64; M];
for a in 0..M {
ex[a] = eta[a].exp();
z += ex[a];
}
let mut p = [0.0_f64; M];
for a in 0..M {
p[a] = ex[a] / z;
}
p
}
fn multinomial_closed_form_vgh<const M: usize>(row: &MultRow<M>) -> KernelChannels<M> {
let p = softmax_active(&row.eta);
let z = {
let mut s = 1.0_f64;
for a in 0..M {
s += row.eta[a].exp();
}
s
};
let obs_eta = if row.obs < M { row.eta[row.obs] } else { 0.0 };
let value = row.w * z.ln() - row.w * obs_eta;
let mut gradient = [0.0_f64; M];
for a in 0..M {
let y_a = if row.obs == a { 1.0 } else { 0.0 };
gradient[a] = row.w * (p[a] - y_a);
}
let mut hessian = [[0.0_f64; M]; M];
for a in 0..M {
for b in 0..M {
let delta = if a == b { 1.0 } else { 0.0 };
hessian[a][b] = row.w * (delta * p[a] - p[a] * p[b]);
}
}
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()
}
}
fn make_rows<const M: usize>(seed: u64, count: usize) -> Vec<MultRow<M>> {
let mut rng = Lcg(seed);
let mut rows = Vec::with_capacity(count);
for i in 0..count {
let mut eta = [0.0_f64; M];
for a in 0..M {
eta[a] = rng.uniform(-2.0, 2.0);
}
rows.push(MultRow {
eta,
obs: i % (M + 1),
w: rng.uniform(0.25, 2.5),
});
}
rows
}
const REL_TOL: f64 = 1e-11;
fn assert_vgh<const M: usize>(seed: u64) {
let rows = make_rows::<M>(seed, 24);
let program = MultinomialSoftmaxRow { rows: rows.clone() };
for (row, fixture) in rows.iter().enumerate() {
let tower: Box<Tower4<M>> =
program_full_tower(&program, row).expect("multinomial jet tower");
let claims = multinomial_closed_form_vgh(fixture);
verify_kernel_channels(&tower, &claims, REL_TOL).unwrap_or_else(|e| {
panic!("M={M} row {row}: softmax closed form disagrees with #932 jet tower: {e}")
});
}
}
#[test]
fn multinomial_softmax_jet_value_grad_hessian_matches_closed_form() {
assert_vgh::<2>(0x9322_2020_1109_face);
assert_vgh::<3>(0x0bad_c0de_2020_1109);
}