Skip to main content

eredu_evaluation/
distribution.rs

1//! Model-distribution comparison independent of backend and modality.
2
3use serde::{Deserialize, Serialize};
4
5/// Metrics comparing one candidate categorical distribution with a reference.
6#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
7pub struct DistributionMetrics {
8    /// KL(reference || candidate), in nats.
9    pub kl_nats: f64,
10    /// Entropy of the reference distribution, in nats.
11    pub reference_entropy_nats: f64,
12    /// Candidate minus reference negative log likelihood for an optional target.
13    pub target_nll_delta_nats: Option<f64>,
14    /// RMSE after removing the mean logit from each side.
15    pub centered_logit_rmse: f64,
16    /// Whether top-1 indices agree.
17    pub top1_agreement: bool,
18    /// Fractional overlap between the selected leading sets.
19    pub top_k_overlap: f64,
20}
21
22/// Compares two complete categorical logit vectors.
23pub 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/// Invalid categorical-distribution comparison.
122#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
123pub enum DistributionError {
124    /// At least one class is required.
125    #[error("distribution must contain at least one class")]
126    Empty,
127    /// Reference and candidate cardinalities differ.
128    #[error("distribution lengths differ: reference {reference}, candidate {candidate}")]
129    Length {
130        /// Reference cardinality.
131        reference: usize,
132        /// Candidate cardinality.
133        candidate: usize,
134    },
135    /// Leading-set size must be positive.
136    #[error("distribution top_k must be positive")]
137    ZeroTopK,
138    /// Target index is outside the distribution.
139    #[error("target index {target} is outside {classes} classes")]
140    Target {
141        /// Invalid target.
142        target: usize,
143        /// Distribution cardinality.
144        classes: usize,
145    },
146    /// Logits must be finite.
147    #[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}