1use crate::error::{invalid_input, require_finite, require_probability};
10use crate::prob::PROB_EPSILON;
11use crate::{ScoringError, ScoringResult};
12
13const MAX_EXACT_F64_INTEGER: u64 = 1u64 << f64::MANTISSA_DIGITS;
14
15#[derive(Debug, Clone, Copy, PartialEq, serde::Serialize, serde::Deserialize)]
32pub struct VectorProbabilityTransform {
33 pub mu_match: f64,
35 pub mu_random: f64,
37 pub sigma: f64,
39 pub base_rate: f64,
41}
42
43impl VectorProbabilityTransform {
44 pub fn new(mu_match: f64, mu_random: f64, sigma: f64, base_rate: f64) -> ScoringResult<Self> {
45 require_finite(mu_match, "mu_match")?;
46 require_finite(mu_random, "mu_random")?;
47 require_finite(sigma, "sigma")?;
48 if sigma <= 0.0 {
49 return Err(invalid_input(format!(
50 "sigma must be positive, got {sigma}"
51 )));
52 }
53 require_probability(base_rate, "base_rate")?;
54 if base_rate == 0.0 || base_rate == 1.0 {
55 return Err(invalid_input(format!(
56 "base_rate must be strictly between 0 and 1, got {base_rate}"
57 )));
58 }
59 Ok(Self {
60 mu_match,
61 mu_random,
62 sigma,
63 base_rate,
64 })
65 }
66
67 pub fn calibrate_one(&self, distance: f64) -> ScoringResult<f64> {
70 let log_lr = self.log_likelihood_ratio(distance)?;
71 let logit_prior = (self.base_rate / (1.0 - self.base_rate)).ln();
72 let logit_post = log_lr + logit_prior;
73 if !logit_post.is_finite() {
74 return Err(ScoringError::ArithmeticOverflow(format!(
75 "calibration logit is not finite for distance {distance}"
76 )));
77 }
78 Ok(1.0 / (1.0 + (-logit_post).exp()))
79 }
80
81 pub fn calibrate(&self, distances: &[f64], weights: Option<&[f64]>) -> ScoringResult<Vec<f64>> {
84 if let Some(weights) = weights {
85 if weights.len() != distances.len() {
86 return Err(invalid_input(format!(
87 "weights length {} does not match distances length {}",
88 weights.len(),
89 distances.len()
90 )));
91 }
92 for (index, weight) in weights.iter().copied().enumerate() {
93 require_finite(weight, &format!("weights[{index}]"))?;
94 }
95 }
96
97 distances
98 .iter()
99 .copied()
100 .enumerate()
101 .map(|(index, distance)| {
102 let posterior = self.calibrate_one(distance)?;
103 let Some(weight) = weights.map(|values| values[index]) else {
104 return Ok(posterior);
105 };
106 let posterior = posterior.clamp(PROB_EPSILON, 1.0 - PROB_EPSILON);
107 let logit = (posterior / (1.0 - posterior)).ln() + weight;
108 if !logit.is_finite() {
109 return Err(ScoringError::ArithmeticOverflow(format!(
110 "weighted calibration logit is not finite at index {index}"
111 )));
112 }
113 Ok(1.0 / (1.0 + (-logit).exp()))
114 })
115 .collect()
116 }
117
118 fn log_likelihood_ratio(&self, distance: f64) -> ScoringResult<f64> {
119 require_finite(distance, "distance")?;
120 let twosq = 2.0 * self.sigma * self.sigma;
124 let r = (self.mu_match - distance).powi(2) / twosq;
125 let g = (self.mu_random - distance).powi(2) / twosq;
126 let ratio = g - r;
129 if ratio.is_finite() {
130 Ok(ratio)
131 } else {
132 Err(ScoringError::ArithmeticOverflow(format!(
133 "likelihood ratio is not finite for distance {distance}"
134 )))
135 }
136 }
137}
138
139pub struct CalibrationMetrics;
140
141#[derive(Debug, Clone, PartialEq)]
142pub struct ReliabilityBin {
143 pub avg_predicted: f64,
144 pub avg_actual: f64,
145 pub count: usize,
146}
147
148#[derive(Debug, Clone, PartialEq)]
149pub struct CalibrationReport {
150 pub ece: f64,
151 pub brier: f64,
152 pub log_loss: f64,
153 pub bins: Vec<ReliabilityBin>,
154}
155
156impl CalibrationMetrics {
157 pub fn log_loss(probabilities: &[f64], labels: &[u8]) -> ScoringResult<f64> {
158 validate_metric_inputs(probabilities, labels)?;
159 if probabilities.is_empty() {
160 return Ok(0.0);
161 }
162 let n = exact_usize_as_f64(probabilities.len(), "probability count")?;
163 let mut s = 0.0;
164 for (&p, &y) in probabilities.iter().zip(labels) {
165 let pp = p.clamp(PROB_EPSILON, 1.0 - PROB_EPSILON);
166 let y = f64::from(y);
167 s += y * pp.ln() + (1.0 - y) * (1.0 - pp).ln();
168 if !s.is_finite() {
169 return Err(ScoringError::ArithmeticOverflow(
170 "log-loss accumulation is not finite".to_string(),
171 ));
172 }
173 }
174 Ok(-s / n)
175 }
176
177 pub fn brier(probabilities: &[f64], labels: &[u8]) -> ScoringResult<f64> {
178 validate_metric_inputs(probabilities, labels)?;
179 if probabilities.is_empty() {
180 return Ok(0.0);
181 }
182 let n = exact_usize_as_f64(probabilities.len(), "probability count")?;
183 let mut sum = 0.0;
184 for (&probability, &label) in probabilities.iter().zip(labels) {
185 sum += (probability - f64::from(label)).powi(2);
186 if !sum.is_finite() {
187 return Err(ScoringError::ArithmeticOverflow(
188 "Brier score accumulation is not finite".to_string(),
189 ));
190 }
191 }
192 Ok(sum / n)
193 }
194
195 pub fn ece(probabilities: &[f64], labels: &[u8], n_bins: usize) -> ScoringResult<f64> {
196 validate_metric_inputs(probabilities, labels)?;
197 validate_bin_count(n_bins)?;
198 let total = probabilities.len();
199 if total == 0 {
200 return Ok(0.0);
201 }
202 let total_f64 = exact_usize_as_f64(total, "probability count")?;
203 let mut acc = 0.0;
204 for bin in reliability_bins(probabilities, labels, n_bins)? {
205 let bin_count = exact_usize_as_f64(bin.count, "reliability bin count")?;
206 acc += (bin_count / total_f64) * (bin.avg_predicted - bin.avg_actual).abs();
207 }
208 Ok(acc)
209 }
210
211 pub fn report(
212 probabilities: &[f64],
213 labels: &[u8],
214 n_bins: usize,
215 ) -> ScoringResult<CalibrationReport> {
216 validate_metric_inputs(probabilities, labels)?;
217 validate_bin_count(n_bins)?;
218 let bins = reliability_bins(probabilities, labels, n_bins)?;
219 let total = exact_usize_as_f64(probabilities.len(), "probability count")?;
220 let ece = if total > 0.0 {
221 let mut sum = 0.0;
222 for bin in &bins {
223 let count = exact_usize_as_f64(bin.count, "reliability bin count")?;
224 sum += (count / total) * (bin.avg_predicted - bin.avg_actual).abs();
225 }
226 sum
227 } else {
228 0.0
229 };
230 Ok(CalibrationReport {
231 ece,
232 brier: Self::brier(probabilities, labels)?,
233 log_loss: Self::log_loss(probabilities, labels)?,
234 bins,
235 })
236 }
237
238 pub fn reliability_diagram(
239 probabilities: &[f64],
240 labels: &[u8],
241 n_bins: usize,
242 ) -> ScoringResult<Vec<ReliabilityBin>> {
243 validate_metric_inputs(probabilities, labels)?;
244 validate_bin_count(n_bins)?;
245 reliability_bins(probabilities, labels, n_bins)
246 }
247}
248
249fn reliability_bins(
250 probabilities: &[f64],
251 labels: &[u8],
252 n_bins: usize,
253) -> ScoringResult<Vec<ReliabilityBin>> {
254 if probabilities.is_empty() {
255 return Ok(Vec::new());
256 }
257 let mut bins: Vec<(f64, f64, usize)> = vec![(0.0, 0.0, 0); n_bins];
258 let n_bins_f = exact_usize_as_f64(n_bins, "reliability bin count")?;
259 for (&p, &y) in probabilities.iter().zip(labels) {
264 let mut idx = (p * n_bins_f) as usize;
265 if idx >= n_bins {
266 idx = n_bins - 1;
267 }
268 if p == 0.0 {
269 idx = 0;
270 }
271 bins[idx].0 += p;
272 bins[idx].1 += f64::from(y);
273 bins[idx].2 = bins[idx].2.checked_add(1).ok_or_else(|| {
274 ScoringError::ArithmeticOverflow("reliability bin count overflow".to_string())
275 })?;
276 }
277 bins.into_iter()
278 .map(|(sum_p, sum_y, count)| -> ScoringResult<_> {
279 if count == 0 {
280 Ok(ReliabilityBin {
281 avg_predicted: 0.0,
282 avg_actual: 0.0,
283 count: 0,
284 })
285 } else {
286 let count_f64 = exact_usize_as_f64(count, "reliability bin count")?;
287 Ok(ReliabilityBin {
288 avg_predicted: sum_p / count_f64,
289 avg_actual: sum_y / count_f64,
290 count,
291 })
292 }
293 })
294 .collect()
295}
296
297fn validate_metric_inputs(probabilities: &[f64], labels: &[u8]) -> ScoringResult<()> {
298 if probabilities.len() != labels.len() {
299 return Err(invalid_input(format!(
300 "probabilities length {} does not match labels length {}",
301 probabilities.len(),
302 labels.len()
303 )));
304 }
305 for (index, probability) in probabilities.iter().copied().enumerate() {
306 require_probability(probability, &format!("probabilities[{index}]"))?;
307 }
308 for (index, label) in labels.iter().copied().enumerate() {
309 if label > 1 {
310 return Err(invalid_input(format!(
311 "labels[{index}] must be 0 or 1, got {label}"
312 )));
313 }
314 }
315 Ok(())
316}
317
318fn validate_bin_count(n_bins: usize) -> ScoringResult<()> {
319 if n_bins == 0 {
320 return Err(invalid_input("n_bins must be greater than zero"));
321 }
322 exact_usize_as_f64(n_bins, "n_bins").map(|_| ())
323}
324
325fn exact_usize_as_f64(value: usize, name: &str) -> ScoringResult<f64> {
326 if u64::try_from(value).is_ok_and(|value| value <= MAX_EXACT_F64_INTEGER) {
327 Ok(value as f64)
328 } else {
329 Err(invalid_input(format!(
330 "{name} {value} exceeds the exact f64 integer range"
331 )))
332 }
333}
334
335#[cfg(test)]
336mod tests {
337 use super::*;
338
339 #[test]
340 fn log_loss_zero_for_perfect_predictions() {
341 let probs = vec![1.0 - PROB_EPSILON, PROB_EPSILON, 1.0 - PROB_EPSILON];
342 let labels = vec![1u8, 0, 1];
343 let loss = CalibrationMetrics::log_loss(&probs, &labels).unwrap();
344 assert!(loss < 1e-8, "expected ~0, got {loss}");
345 }
346
347 #[test]
348 fn brier_zero_for_perfect_predictions() {
349 let probs = vec![1.0, 0.0, 1.0];
350 let labels = vec![1u8, 0, 1];
351 let brier = CalibrationMetrics::brier(&probs, &labels).unwrap();
352 assert!(brier < 1e-12);
353 }
354
355 #[test]
356 fn ece_zero_when_perfectly_calibrated() {
357 let probs = vec![0.05; 100]; let mut labels = vec![0u8; 100];
360 for label in &mut labels[..5] {
361 *label = 1;
362 }
363 let ece = CalibrationMetrics::ece(&probs, &labels, 10).unwrap();
364 assert!(ece < 1e-9, "got {ece}");
365 }
366
367 #[test]
368 fn transform_and_metrics_reject_invalid_inputs() {
369 assert!(VectorProbabilityTransform::new(0.0, 1.0, 0.0, 0.5).is_err());
370 assert!(VectorProbabilityTransform::new(0.0, 1.0, 1.0, 1.0).is_err());
371 let transform = VectorProbabilityTransform::new(0.0, 1.0, 1.0, 0.5).unwrap();
372 assert!(transform.calibrate_one(f64::NAN).is_err());
373 assert!(transform.calibrate(&[0.1], Some(&[])).is_err());
374
375 assert!(CalibrationMetrics::log_loss(&[0.5], &[]).is_err());
376 assert!(CalibrationMetrics::brier(&[f64::NAN], &[0]).is_err());
377 assert!(CalibrationMetrics::ece(&[0.5], &[2], 10).is_err());
378 assert!(CalibrationMetrics::ece(&[0.5], &[0], 0).is_err());
379 }
380}