eredu_evaluation/
distribution.rs1use serde::{Deserialize, Serialize};
4
5#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
7pub struct DistributionMetrics {
8 pub kl_nats: f64,
10 pub reference_entropy_nats: f64,
12 pub target_nll_delta_nats: Option<f64>,
14 pub centered_logit_rmse: f64,
16 pub top1_agreement: bool,
18 pub top_k_overlap: f64,
20}
21
22pub fn compare_distributions(
24 reference: &[f32],
25 candidate: &[f32],
26 target: Option<usize>,
27 top_k: usize,
28) -> Result<DistributionMetrics, DistributionError> {
29 if reference.is_empty() {
30 return Err(DistributionError::Empty);
31 }
32 if reference.len() != candidate.len() {
33 return Err(DistributionError::Length {
34 reference: reference.len(),
35 candidate: candidate.len(),
36 });
37 }
38 if top_k == 0 {
39 return Err(DistributionError::ZeroTopK);
40 }
41 if let Some(target) = target {
42 if target >= reference.len() {
43 return Err(DistributionError::Target {
44 target,
45 classes: reference.len(),
46 });
47 }
48 }
49 if reference
50 .iter()
51 .chain(candidate)
52 .any(|value| !value.is_finite())
53 {
54 return Err(DistributionError::NonFinite);
55 }
56
57 let reference_lse = logsumexp(reference);
58 let candidate_lse = logsumexp(candidate);
59 let mut kl = 0.0;
60 let mut entropy = 0.0;
61 for (&reference_logit, &candidate_logit) in reference.iter().zip(candidate) {
62 let log_p = reference_logit as f64 - reference_lse;
63 let log_q = candidate_logit as f64 - candidate_lse;
64 let probability = log_p.exp();
65 kl += probability * (log_p - log_q);
66 entropy -= probability * log_p;
67 }
68 let reference_mean =
69 reference.iter().map(|value| *value as f64).sum::<f64>() / reference.len() as f64;
70 let candidate_mean =
71 candidate.iter().map(|value| *value as f64).sum::<f64>() / candidate.len() as f64;
72 let centered_logit_rmse = (reference
73 .iter()
74 .zip(candidate)
75 .map(|(&left, &right)| {
76 let delta = (left as f64 - reference_mean) - (right as f64 - candidate_mean);
77 delta * delta
78 })
79 .sum::<f64>()
80 / reference.len() as f64)
81 .sqrt();
82 let reference_top = top_indices(reference, top_k);
83 let candidate_top = top_indices(candidate, top_k);
84 let overlap = reference_top
85 .iter()
86 .filter(|index| candidate_top.contains(index))
87 .count();
88 Ok(DistributionMetrics {
89 kl_nats: kl.max(0.0),
90 reference_entropy_nats: entropy,
91 target_nll_delta_nats: target.map(|target| {
92 (candidate_lse - candidate[target] as f64) - (reference_lse - reference[target] as f64)
93 }),
94 centered_logit_rmse,
95 top1_agreement: reference_top[0] == candidate_top[0],
96 top_k_overlap: overlap as f64 / reference_top.len() as f64,
97 })
98}
99
100fn logsumexp(values: &[f32]) -> f64 {
101 let maximum = values.iter().copied().fold(f32::NEG_INFINITY, f32::max) as f64;
102 maximum
103 + values
104 .iter()
105 .map(|value| (*value as f64 - maximum).exp())
106 .sum::<f64>()
107 .ln()
108}
109
110fn top_indices(values: &[f32], count: usize) -> Vec<usize> {
111 let mut indices = (0..values.len()).collect::<Vec<_>>();
112 indices.sort_by(|left, right| {
113 values[*right]
114 .total_cmp(&values[*left])
115 .then_with(|| left.cmp(right))
116 });
117 indices.truncate(count.min(indices.len()));
118 indices
119}
120
121#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
123pub enum DistributionError {
124 #[error("distribution must contain at least one class")]
126 Empty,
127 #[error("distribution lengths differ: reference {reference}, candidate {candidate}")]
129 Length {
130 reference: usize,
132 candidate: usize,
134 },
135 #[error("distribution top_k must be positive")]
137 ZeroTopK,
138 #[error("target index {target} is outside {classes} classes")]
140 Target {
141 target: usize,
143 classes: usize,
145 },
146 #[error("distribution logits must be finite")]
148 NonFinite,
149}
150
151#[cfg(test)]
152mod tests {
153 use super::*;
154
155 #[test]
156 fn identical_distributions_have_exact_metrics() {
157 let metrics =
158 compare_distributions(&[0.0, 2.0, 1.0], &[0.0, 2.0, 1.0], Some(1), 2).unwrap();
159 assert_eq!(metrics.kl_nats, 0.0);
160 assert_eq!(metrics.target_nll_delta_nats, Some(0.0));
161 assert_eq!(metrics.centered_logit_rmse, 0.0);
162 assert!(metrics.top1_agreement);
163 assert_eq!(metrics.top_k_overlap, 1.0);
164 }
165}