pub(super) fn softmax(logits: &[f32]) -> Vec<f64> {
let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64;
let mut out: Vec<f64> = logits.iter().map(|&v| (v as f64 - max).exp()).collect();
let sum: f64 = out.iter().sum();
if sum > 0.0 {
for v in &mut out {
*v /= sum;
}
}
out
}
pub(super) fn kl(p: &[f64], q: &[f64]) -> f64 {
let mut acc = 0.0f64;
for (&pi, &qi) in p.iter().zip(q) {
if pi > 0.0 {
acc += pi * (pi.max(f64::MIN_POSITIVE) / qi.max(f64::MIN_POSITIVE)).ln();
}
}
acc
}
pub(super) fn total_variation(p: &[f64], q: &[f64]) -> (f64, f64) {
let mut tv = 0.0f64;
let mut max_delta = 0.0f64;
for (&pi, &qi) in p.iter().zip(q) {
let d = (pi - qi).abs();
tv += d;
if d > max_delta {
max_delta = d;
}
}
(0.5 * tv, max_delta)
}
pub(super) fn order_desc(p: &[f64]) -> Vec<usize> {
let mut idx: Vec<usize> = (0..p.len()).collect();
idx.sort_by(|&a, &b| p[b].partial_cmp(&p[a]).unwrap().then(a.cmp(&b)));
idx
}
pub(super) fn ulps_between(a: f32, b: f32) -> i64 {
fn key(x: f32) -> i64 {
let bits = i64::from(x.to_bits() & 0x7fff_ffff);
if x.is_sign_negative() {
-bits
} else {
bits
}
}
(key(a) - key(b)).abs()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ulps_between_counts_representable_values_not_absolute_distance() {
for anchor in [3.0f32, 3e-8, -3.0, 1.0, 65_536.0] {
let next = f32::from_bits(if anchor.is_sign_negative() {
anchor.to_bits() - 1
} else {
anchor.to_bits() + 1
});
assert_eq!(
ulps_between(anchor, next),
1,
"{anchor} and its neighbour {next} must be one ulp apart"
);
}
assert_eq!(ulps_between(1.5, 1.5), 0);
assert_eq!(ulps_between(2.0, 1.0), ulps_between(1.0, 2.0));
assert_eq!(ulps_between(-0.0, 0.0), 0);
let tiny = f32::from_bits(1);
assert_eq!(ulps_between(-tiny, tiny), 2);
assert_eq!(ulps_between(1.0, 2.0), 1 << 23);
}
#[test]
fn kl_is_zero_for_identical_distributions_and_asymmetric_otherwise() {
let p = softmax(&[0.1f32, 5.0, -2.0, 3.3]);
assert_eq!(kl(&p, &p), 0.0);
let q = softmax(&[0.1f32, 5.0, -2.0, 1.0]);
assert!(kl(&p, &q) > 0.0 && kl(&q, &p) > 0.0);
assert_ne!(
kl(&p, &q),
kl(&q, &p),
"if these ever agree the spread's `max over ordered pairs` is pointless"
);
let with_a_dead_token = softmax(&[0.0f32, -800.0]);
assert_eq!(
with_a_dead_token[1], 0.0,
"the fixture must actually contain an exact zero"
);
let uniform = softmax(&[0.0f32, 0.0]);
let d = kl(&with_a_dead_token, &uniform);
assert!(d.is_finite(), "a zero-mass token made the KL {d}");
assert!((d - std::f64::consts::LN_2).abs() < 1e-12, "{d}");
}
#[test]
fn softmax_is_shift_invariant_even_where_exp_would_overflow() {
let base = [0.5f32, 1.5, -3.0];
let a = softmax(&base);
let shifted: Vec<f32> = base.iter().map(|v| v + 800.0).collect();
let b = softmax(&shifted);
for (x, y) in a.iter().zip(&b) {
assert!(x.is_finite() && y.is_finite(), "{x} vs {y}");
assert!((x - y).abs() < 1e-12, "{x} vs {y}");
}
assert!((a.iter().sum::<f64>() - 1.0).abs() < 1e-12);
assert!((b.iter().sum::<f64>() - 1.0).abs() < 1e-12);
}
}