Skip to main content

uqa_scoring/
calibration.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Calibration diagnostics (Paper 3 Section 11.3, Paper 5 Section 8.3).
8
9use 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/// Likelihood-ratio calibrator for vector distances (Theorem 3.1.1,
16/// Paper 5). Converts vector similarity into calibrated probability.
17///
18/// The transform models distances to relevant documents (`f_R`) and
19/// to background (random) documents (`f_G`) as Gaussian distributions.
20/// Calibration converts a distance into a posterior probability by
21/// applying Bayes' rule with the configured base rate:
22///
23/// ```text
24/// log f_R(d) - log f_G(d) + logit(base_rate)  ->  sigmoid
25/// ```
26///
27/// The formulation is deliberately small; downstream callers fit the
28/// means and standard deviations offline (for example, via the parameter
29/// learner) and pass the transform through. Optional per-distance
30/// weights bias the computed log-odds before the sigmoid.
31#[derive(Debug, Clone, Copy, PartialEq, serde::Serialize, serde::Deserialize)]
32pub struct VectorProbabilityTransform {
33    /// Mean distance for relevant documents (numerator distribution).
34    pub mu_match: f64,
35    /// Mean distance for background documents (denominator distribution).
36    pub mu_random: f64,
37    /// Shared standard deviation. Must be positive.
38    pub sigma: f64,
39    /// Prior probability of relevance. Default 0.5 (neutral).
40    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    /// Convert a single distance to a probability via the likelihood
68    /// ratio + base-rate logit.
69    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    /// Vectorized calibration. Optional `weights` bias each distance's
82    /// log-odds before the sigmoid.
83    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        // Gaussian log-LR with shared sigma:
121        // (mu_R - d)^2 / (2 sigma^2) - (mu_G - d)^2 / (2 sigma^2)
122        // collapses to a linear function of d.
123        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        // Numerator should DOMINATE for small (near-relevant) distances,
127        // so the log-LR is `g - r`.
128        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    // Bin boundaries match the upstream UQA ECE contract: lowest bin is
260    // `[lo, hi]` (inclusive both ends), the rest are `(lo, hi]`. Floor
261    // division places exact upper edges into the higher bin, except for
262    // `p == 0.0` which belongs to bin 0.
263    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        // Each bin's avg predicted == avg actual.
358        let probs = vec![0.05; 100]; // 5% predicted
359        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}