Skip to main content

cerno_core/
math.rs

1//! Turning label logprobs into probabilities and a confidence.
2
3/// Softmax over `logprobs`, with `temperature` scaling applied first.
4///
5/// Scaling before the exponential is plain temperature scaling: `T = 1` reproduces the model's
6/// own distribution, `T > 1` flattens it. Small instruct-tuned models answer clear cases at a
7/// probability of 1.0000, so flattening is usually what a caller wants — but it is their call,
8/// which is why the raw logprobs travel back in every answer.
9///
10/// Subtracting the maximum before exponentiating keeps the sum finite for the very negative
11/// logprobs a truncated distribution produces (values around `-25` are routine).
12pub fn softmax(logprobs: &[f64], temperature: f64) -> Vec<f64> {
13    if logprobs.is_empty() {
14        return Vec::new();
15    }
16
17    let scaled: Vec<f64> = logprobs.iter().map(|lp| lp / temperature).collect();
18    let max = scaled.iter().copied().fold(f64::NEG_INFINITY, f64::max);
19    let exponentiated: Vec<f64> = scaled.iter().map(|s| (s - max).exp()).collect();
20    let sum: f64 = exponentiated.iter().sum();
21
22    if sum <= 0.0 || !sum.is_finite() {
23        // Every candidate underflowed. Uniform is the only honest answer.
24        let uniform = 1.0 / logprobs.len() as f64;
25        return vec![uniform; logprobs.len()];
26    }
27
28    exponentiated.iter().map(|e| e / sum).collect()
29}
30
31/// How peaked a distribution is, in `0.0..=1.0`.
32///
33/// This is `1 - H(p)/log(n)`: one when all mass sits on a single label, zero when the labels are
34/// indistinguishable. Dividing by `log(n)` puts every primitive on the same `0..=1` scale, so a
35/// two-option Noul and a ten-level Score can be thresholded by the same rule — one definition for
36/// all three beats three special cases no caller can hold in their head at once.
37///
38/// It is a scale, not an invariant. At a fixed top probability the value *rises* with the number
39/// of labels: 0.9 on one of ten options scores 0.76, 0.9 on one of two scores 0.53. That is the
40/// measure working as intended — narrowing ten candidates down to one is the stronger statement —
41/// but it does mean confidence values are only directly comparable between questions of the same
42/// shape.
43pub fn confidence(probabilities: &[f64]) -> f64 {
44    let n = probabilities.len();
45    if n < 2 {
46        return 1.0;
47    }
48
49    let entropy: f64 = probabilities
50        .iter()
51        .filter(|p| **p > 0.0)
52        .map(|p| -p * p.ln())
53        .sum();
54
55    (1.0 - entropy / (n as f64).ln()).clamp(0.0, 1.0)
56}
57
58/// The probability-weighted mean of 1-based level indices.
59pub fn expected_level(probabilities: &[f64]) -> f64 {
60    probabilities
61        .iter()
62        .enumerate()
63        .map(|(i, p)| (i + 1) as f64 * p)
64        .sum()
65}
66
67/// The index of the largest probability. Ties go to the lower index.
68pub fn argmax(probabilities: &[f64]) -> usize {
69    probabilities
70        .iter()
71        .enumerate()
72        .fold((0usize, f64::NEG_INFINITY), |(bi, bv), (i, v)| {
73            if *v > bv { (i, *v) } else { (bi, bv) }
74        })
75        .0
76}
77
78/// Combine logprobs of several tokens that mean the same answer, in probability space.
79///
80/// Tokenisers spell one answer many ways. Measured on `gemma4:26b`, asking for a yes/no in
81/// German surfaced `Ja`, `JA`, ` Ja` and `ja` as four separate entries, each holding part of the
82/// mass — and `Nein` was not a token at all, arriving as `Ne` + `in`. Single-letter labels avoid
83/// the split entirely, but a leading-space variant such as `" A"` still shows up on some
84/// tokenisers, so every variant of a label is folded together here rather than one of them being
85/// picked and the rest discarded.
86pub fn logsumexp(logprobs: &[f64]) -> f64 {
87    if logprobs.is_empty() {
88        return f64::NEG_INFINITY;
89    }
90
91    let max = logprobs.iter().copied().fold(f64::NEG_INFINITY, f64::max);
92    if !max.is_finite() {
93        return max;
94    }
95
96    max + logprobs.iter().map(|lp| (lp - max).exp()).sum::<f64>().ln()
97}
98
99#[cfg(test)]
100mod tests {
101    use super::*;
102
103    fn close(a: f64, b: f64) -> bool {
104        (a - b).abs() < 1e-9
105    }
106
107    #[test]
108    fn softmax_of_equal_logprobs_is_uniform() {
109        let p = softmax(&[-1.0, -1.0, -1.0, -1.0], 1.0);
110
111        assert!(p.iter().all(|v| close(*v, 0.25)));
112    }
113
114    /// Hand-checked against the measured `gemma4:26b` distribution: with logprobs -0.005 and
115    /// -5.246 the difference is 5.241 nats, so the odds are e^5.241 : 1.
116    #[test]
117    fn softmax_matches_a_hand_computed_pair() {
118        let p = softmax(&[-0.005, -5.246], 1.0);
119
120        let expected_top = 1.0 / (1.0 + (-5.241f64).exp());
121        assert!(close(p[0], expected_top), "got {}", p[0]);
122        assert!(close(p[0] + p[1], 1.0));
123    }
124
125    #[test]
126    fn softmax_sums_to_one() {
127        let p = softmax(&[-0.005, -5.246, -11.515, -13.662], 1.0);
128
129        assert!(close(p.iter().sum::<f64>(), 1.0));
130    }
131
132    /// The point of calibration: same ranking, less certainty.
133    #[test]
134    fn higher_temperature_flattens_without_reordering() {
135        let raw = softmax(&[-0.005, -5.246, -11.515], 1.0);
136        let warm = softmax(&[-0.005, -5.246, -11.515], 3.0);
137
138        assert!(warm[0] < raw[0], "{} should be below {}", warm[0], raw[0]);
139        assert!(warm[1] > raw[1]);
140        assert_eq!(argmax(&raw), argmax(&warm));
141        assert!(close(warm.iter().sum::<f64>(), 1.0));
142    }
143
144    /// Extremely negative logprobs are the normal case once a label falls back to the floor.
145    #[test]
146    fn extreme_logprobs_do_not_produce_nan() {
147        let p = softmax(&[0.0, -800.0, -1500.0], 1.0);
148
149        assert!(p.iter().all(|v| v.is_finite()));
150        assert!(close(p.iter().sum::<f64>(), 1.0));
151        assert!(close(p[0], 1.0));
152    }
153
154    #[test]
155    fn confidence_is_one_when_all_mass_is_on_one_label() {
156        assert!(close(confidence(&[1.0, 0.0, 0.0]), 1.0));
157    }
158
159    #[test]
160    fn confidence_is_zero_for_a_uniform_distribution() {
161        assert!(close(confidence(&[0.25; 4]), 0.0));
162        assert!(close(confidence(&[0.5, 0.5]), 0.0));
163    }
164
165    /// Confidence rises as mass concentrates, for a fixed number of labels.
166    #[test]
167    fn confidence_rises_with_concentration() {
168        let flat = confidence(&[0.4, 0.3, 0.3]);
169        let peaked = confidence(&[0.8, 0.1, 0.1]);
170        let certain = confidence(&[0.98, 0.01, 0.01]);
171
172        assert!(flat < peaked, "{flat} should be below {peaked}");
173        assert!(peaked < certain, "{peaked} should be below {certain}");
174    }
175
176    /// Every primitive lands on the same 0..=1 scale, whatever its label count.
177    #[test]
178    fn confidence_stays_within_the_unit_range() {
179        for n in 2..=10 {
180            let mut p = vec![0.02 / (n - 1) as f64; n];
181            p[0] = 0.98;
182
183            let c = confidence(&p);
184            assert!((0.0..=1.0).contains(&c), "n={n} gave {c}");
185        }
186    }
187
188    /// The same top probability scores *higher* with more labels, because picking one of ten is
189    /// the stronger statement. Pinned with hand-computed values so the semantics stay documented.
190    #[test]
191    fn confidence_is_not_invariant_to_label_count() {
192        let two = confidence(&[0.9, 0.1]);
193        let ten = {
194            let mut p = vec![0.1 / 9.0; 10];
195            p[0] = 0.9;
196            confidence(&p)
197        };
198
199        assert!(close(two, 0.531_004_406_410_718_9), "got {two}");
200        assert!(close(ten, 0.763_394_007_551_459_9), "got {ten}");
201        assert!(two < ten);
202    }
203
204    #[test]
205    fn expected_level_weights_by_probability() {
206        // Split evenly between level 1 and level 5.
207        assert!(close(expected_level(&[0.5, 0.0, 0.0, 0.0, 0.5]), 3.0));
208        // All mass on level 4.
209        assert!(close(expected_level(&[0.0, 0.0, 0.0, 1.0, 0.0]), 4.0));
210    }
211
212    #[test]
213    fn argmax_picks_the_largest_and_breaks_ties_low() {
214        assert_eq!(argmax(&[0.1, 0.7, 0.2]), 1);
215        assert_eq!(argmax(&[0.5, 0.5]), 0);
216    }
217
218    /// Two tokens at equal probability must fold to exactly twice the mass.
219    #[test]
220    fn logsumexp_adds_probability_mass() {
221        let combined = logsumexp(&[-1.0, -1.0]);
222
223        assert!(
224            close(combined.exp(), 2.0 * (-1.0f64).exp()),
225            "got {}",
226            combined.exp()
227        );
228    }
229
230    #[test]
231    fn logsumexp_of_a_single_value_is_that_value() {
232        assert!(close(logsumexp(&[-3.25]), -3.25));
233    }
234
235    #[test]
236    fn logsumexp_of_nothing_is_negative_infinity() {
237        assert_eq!(logsumexp(&[]), f64::NEG_INFINITY);
238    }
239
240    #[test]
241    fn logsumexp_is_dominated_by_the_largest_term() {
242        let combined = logsumexp(&[-0.001, -30.0]);
243
244        assert!((combined - -0.001).abs() < 1e-6, "got {combined}");
245    }
246}