pub const TOPK: usize = 10;
pub fn log_softmax_f64(logits: &[f32]) -> Vec<f64> {
let mut max = f64::NEG_INFINITY;
for &l in logits {
let l = f64::from(l);
if l > max {
max = l;
}
}
let mut sum = 0.0f64;
let mut out = Vec::with_capacity(logits.len());
for &l in logits {
let d = f64::from(l) - max;
sum += d.exp();
out.push(d);
}
let lse = sum.ln();
for v in &mut out {
*v -= lse;
}
out
}
pub fn topk_indices(logits: &[f32], k: usize) -> Vec<u32> {
let mut idx: Vec<u32> = (0..logits.len() as u32).collect();
let k = k.min(idx.len());
idx.sort_unstable_by(|&a, &b| {
logits[b as usize]
.partial_cmp(&logits[a as usize])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
idx.truncate(k);
idx
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct StepCompare {
pub kld: f64,
pub top1_agree: bool,
pub topk_overlap: f64,
pub ref_logprob_target: f64,
pub arm_logprob_target: f64,
}
pub fn compare_step(ref_logits: &[f32], arm_logits: &[f32], target: u32) -> StepCompare {
assert_eq!(
ref_logits.len(),
arm_logits.len(),
"vocab mismatch between reference and arm"
);
let t = target as usize;
assert!(t < ref_logits.len(), "target {target} outside vocab");
let lp = log_softmax_f64(ref_logits);
let lq = log_softmax_f64(arm_logits);
let mut kld = 0.0f64;
for (a, b) in lp.iter().zip(&lq) {
let p = a.exp();
if p > 0.0 {
kld += p * (a - b);
}
}
let ref_top = topk_indices(ref_logits, TOPK);
let arm_top = topk_indices(arm_logits, TOPK);
let hits = ref_top.iter().filter(|i| arm_top.contains(i)).count();
StepCompare {
kld,
top1_agree: ref_top[0] == arm_top[0],
topk_overlap: hits as f64 / ref_top.len() as f64,
ref_logprob_target: lp[t],
arm_logprob_target: lq[t],
}
}
pub fn kld_reverse(ref_logits: &[f32], arm_logits: &[f32]) -> f64 {
compare_step(arm_logits, ref_logits, 0).kld
}
#[derive(Clone, Debug, Default)]
pub struct KldAccum {
klds: Vec<f64>,
top1_hits: usize,
topk_overlap_sum: f64,
ref_nll: f64,
arm_nll: f64,
}
impl KldAccum {
pub fn new() -> Self {
Self::default()
}
pub fn push(&mut self, c: StepCompare) {
self.klds.push(c.kld);
self.top1_hits += usize::from(c.top1_agree);
self.topk_overlap_sum += c.topk_overlap;
self.ref_nll -= c.ref_logprob_target;
self.arm_nll -= c.arm_logprob_target;
}
pub fn len(&self) -> usize {
self.klds.len()
}
pub fn is_empty(&self) -> bool {
self.klds.is_empty()
}
pub fn summary(&self) -> Option<KldSummary> {
if self.klds.is_empty() {
return None;
}
let n = self.klds.len();
let mut sorted = self.klds.clone();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
Some(KldSummary {
n,
kld_mean: self.klds.iter().sum::<f64>() / n as f64,
kld_median: quantile(&sorted, 0.5),
kld_p99: quantile(&sorted, 0.99),
kld_max: sorted[n - 1],
top1_agreement: self.top1_hits as f64 / n as f64,
topk_overlap: self.topk_overlap_sum / n as f64,
ppl_ref: (self.ref_nll / n as f64).exp(),
ppl_arm: (self.arm_nll / n as f64).exp(),
})
}
}
fn quantile(sorted: &[f64], q: f64) -> f64 {
let n = sorted.len();
let rank = (q * n as f64).ceil() as usize;
sorted[rank.clamp(1, n) - 1]
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct KldSummary {
pub n: usize,
pub kld_mean: f64,
pub kld_median: f64,
pub kld_p99: f64,
pub kld_max: f64,
pub top1_agreement: f64,
pub topk_overlap: f64,
pub ppl_ref: f64,
pub ppl_arm: f64,
}
#[cfg(test)]
mod tests {
use super::*;
fn logits2(p0: f64, p1: f64) -> Vec<f32> {
vec![p0.ln() as f32, p1.ln() as f32]
}
#[test]
fn self_vs_self_is_exactly_zero() {
let l: Vec<f32> = (0..1024).map(|i| ((i * 37 % 101) as f32) * 0.1 - 5.0).collect();
let c = compare_step(&l, &l, 7);
assert_eq!(c.kld, 0.0, "self-KLD must be exactly zero");
assert!(c.top1_agree);
assert_eq!(c.topk_overlap, 1.0);
assert_eq!(c.ref_logprob_target, c.arm_logprob_target);
}
#[test]
fn matches_hand_computed_two_point_kld() {
let c = compare_step(&logits2(0.5, 0.5), &logits2(0.75, 0.25), 0);
assert!(
(c.kld - 0.143_841_0).abs() < 1e-6,
"hand-computed KLD mismatch: {}",
c.kld
);
}
#[test]
fn kld_is_asymmetric_and_reverse_agrees() {
let (p, q) = (logits2(0.5, 0.5), logits2(0.75, 0.25));
let fwd = compare_step(&p, &q, 0).kld;
let rev = kld_reverse(&p, &q);
assert!((rev - 0.130_812_0).abs() < 1e-6, "reverse KLD: {rev}");
assert!(
(fwd - rev).abs() > 1e-3,
"the two directions must not be conflated"
);
}
#[test]
fn kld_is_shift_invariant() {
let a: Vec<f32> = vec![1.0, 2.0, 3.0, 0.5];
let b: Vec<f32> = vec![1.5, 2.0, 2.5, 0.0];
let base = compare_step(&a, &b, 2).kld;
assert!(base > 1e-3, "base case must be a real divergence: {base}");
let sa: Vec<f32> = a.iter().map(|v| v + 800.0).collect();
let sb: Vec<f32> = b.iter().map(|v| v + 800.0).collect();
let shifted = compare_step(&sa, &sb, 2).kld;
assert!(base.is_finite() && shifted.is_finite());
assert!(
(base - shifted).abs() < 1e-12,
"shift changed KLD: {base} vs {shifted}"
);
}
#[test]
fn perturbation_raises_kld_and_drops_agreement() {
let refl: Vec<f32> = (0..512).map(|i| ((i % 17) as f32) * 0.3).collect();
let mut arm = refl.clone();
arm[3] += 9.0; for (i, v) in arm.iter_mut().enumerate() {
*v += ((i % 5) as f32) * 0.05; }
let c = compare_step(&refl, &arm, 11);
assert!(c.kld > 0.0, "perturbation must show positive KLD");
assert!(!c.top1_agree, "argmax should have moved");
assert!(c.topk_overlap < 1.0, "top-k set should have moved");
}
#[test]
fn kld_is_non_negative_on_random_pairs() {
let mut seed = 0x2545_F491_4F6C_DD1Du64;
let mut next = || {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
(seed >> 40) as f32 / 1024.0 - 12.0
};
for _ in 0..64 {
let a: Vec<f32> = (0..256).map(|_| next()).collect();
let b: Vec<f32> = (0..256).map(|_| next()).collect();
let k = compare_step(&a, &b, 0).kld;
assert!(k >= -1e-9, "KLD went negative: {k}");
}
}
#[test]
fn topk_overlap_counts_set_intersection_not_order() {
let a: Vec<f32> = vec![5.0, 4.0, 3.0, 0.0, 0.0];
let b: Vec<f32> = vec![3.0, 4.0, 5.0, 0.0, 0.0];
assert_eq!(topk_indices(&a, 3), vec![0, 1, 2]);
assert_eq!(topk_indices(&b, 3), vec![2, 1, 0]);
let hits = topk_indices(&a, 3)
.iter()
.filter(|i| topk_indices(&b, 3).contains(i))
.count();
assert_eq!(hits, 3);
}
#[test]
fn topk_ties_break_by_index_deterministically() {
let flat: Vec<f32> = vec![1.0; 8];
assert_eq!(topk_indices(&flat, 3), vec![0, 1, 2]);
}
#[test]
fn uniform_distribution_has_perplexity_equal_to_vocab() {
let flat: Vec<f32> = vec![0.0; 64];
let mut acc = KldAccum::new();
for t in 0..10u32 {
acc.push(compare_step(&flat, &flat, t));
}
let s = acc.summary().expect("non-empty");
assert!((s.ppl_ref - 64.0).abs() < 1e-9, "ppl_ref {}", s.ppl_ref);
assert!((s.ppl_arm - 64.0).abs() < 1e-9, "ppl_arm {}", s.ppl_arm);
assert_eq!(s.kld_mean, 0.0);
assert_eq!(s.top1_agreement, 1.0);
assert_eq!(s.n, 10);
}
#[test]
fn summary_order_statistics_are_nearest_rank() {
let mut acc = KldAccum::new();
for i in 0..100 {
acc.push(StepCompare {
kld: i as f64 / 100.0,
top1_agree: i < 90,
topk_overlap: 1.0,
ref_logprob_target: -1.0,
arm_logprob_target: -1.0,
});
}
let s = acc.summary().expect("non-empty");
assert_eq!(s.n, 100);
assert!((s.kld_median - 0.49).abs() < 1e-12, "median {}", s.kld_median);
assert!((s.kld_p99 - 0.98).abs() < 1e-12, "p99 {}", s.kld_p99);
assert!((s.kld_max - 0.99).abs() < 1e-12);
assert!((s.top1_agreement - 0.90).abs() < 1e-12);
assert!((s.ppl_ref - std::f64::consts::E).abs() < 1e-12);
}
#[test]
fn quantile_uses_nearest_rank_not_floor() {
let v: Vec<f64> = (0..7).map(|i| i as f64).collect();
assert_eq!(quantile(&v, 0.5), 3.0, "rank ceil(3.5) = 4th smallest");
assert_eq!(quantile(&v, 0.99), 6.0, "rank ceil(6.93) = 7th smallest");
assert_eq!(quantile(&v, 1.0), 6.0);
assert_eq!(quantile(&v, 0.0), 0.0, "rank clamps up to 1");
}
#[test]
fn empty_accumulator_yields_no_summary() {
assert!(KldAccum::new().summary().is_none());
}
}