1pub 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 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
31pub 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
58pub 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
67pub 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
78pub 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 #[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 #[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 #[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 #[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 #[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 #[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 assert!(close(expected_level(&[0.5, 0.0, 0.0, 0.0, 0.5]), 3.0));
208 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 #[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}