Skip to main content

uqa_scoring/
parameter_learner.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Supervised fitting for query-level Bayesian BM25 calibration.
8//!
9//! The fitted posterior is the same transform used by
10//! `BayesianBM25Scorer`: `sigmoid(alpha * (raw_bm25_score - beta))`.
11//! Any intercept the labels call for is absorbed into `beta`. The
12//! `base_rate` is not a model term: it tracks the observed positive
13//! label rate as an exponential moving average, estimating the corpus
14//! relevance prior that fusion applies exactly once.
15
16use std::collections::BTreeMap;
17
18use crate::error::{invalid_input, require_finite, require_probability};
19use crate::prob::sigmoid;
20use crate::{ScoringError, ScoringResult};
21
22#[derive(Debug, Clone)]
23pub struct ParameterLearner {
24    alpha: f64,
25    beta: f64,
26    base_rate: Option<f64>,
27}
28
29impl ParameterLearner {
30    pub fn new(alpha: f64, beta: f64, base_rate: Option<f64>) -> ScoringResult<Self> {
31        require_finite(alpha, "alpha")?;
32        if alpha <= 0.0 {
33            return Err(invalid_input(format!(
34                "alpha must be positive, got {alpha}"
35            )));
36        }
37        require_finite(beta, "beta")?;
38        if let Some(base_rate) = base_rate {
39            require_probability(base_rate, "base_rate")?;
40            if base_rate == 0.0 || base_rate == 1.0 {
41                return Err(invalid_input(format!(
42                    "base_rate must be strictly between 0 and 1, got {base_rate}"
43                )));
44            }
45        }
46        Ok(Self {
47            alpha,
48            beta,
49            base_rate,
50        })
51    }
52
53    pub fn alpha(&self) -> f64 {
54        self.alpha
55    }
56
57    pub fn beta(&self) -> f64 {
58        self.beta
59    }
60
61    pub fn base_rate(&self) -> Option<f64> {
62        self.base_rate
63    }
64
65    pub fn params(&self) -> BTreeMap<String, f64> {
66        BTreeMap::from([
67            ("alpha".to_string(), self.alpha),
68            ("beta".to_string(), self.beta),
69            ("base_rate".to_string(), self.base_rate.unwrap_or(0.0)),
70        ])
71    }
72
73    pub fn probability(&self, raw_score: f64) -> ScoringResult<f64> {
74        require_finite(raw_score, "raw_score")?;
75        let centered = raw_score - self.beta;
76        let logit = self.alpha * centered;
77        if !centered.is_finite() || !logit.is_finite() {
78            return Err(ScoringError::ArithmeticOverflow(format!(
79                "learner logit is not finite for score {raw_score}"
80            )));
81        }
82        Ok(sigmoid(logit))
83    }
84
85    /// Apply one exact logistic-loss gradient step to a raw BM25 score.
86    pub fn update(&mut self, raw_score: f64, label: f64, learning_rate: f64) -> ScoringResult<()> {
87        require_finite(raw_score, "raw_score")?;
88        require_probability(label, "label")?;
89        require_finite(learning_rate, "learning_rate")?;
90        if learning_rate <= 0.0 {
91            return Err(invalid_input(format!(
92                "learning_rate must be positive, got {learning_rate}"
93            )));
94        }
95
96        let error = self.probability(raw_score)? - label;
97        let alpha_before = self.alpha;
98        let beta_before = self.beta;
99        let next_alpha = self.alpha - learning_rate * error * (raw_score - beta_before);
100        let next_beta = self.beta - learning_rate * error * -alpha_before;
101        if !next_alpha.is_finite() || next_alpha <= 0.0 || !next_beta.is_finite() {
102            return Err(ScoringError::ArithmeticOverflow(format!(
103                "gradient update produced invalid parameters alpha={next_alpha}, beta={next_beta}"
104            )));
105        }
106
107        let next_base_rate = if let Some(base_rate) = self.base_rate {
108            let updated = base_rate + learning_rate * (label - base_rate);
109            if !updated.is_finite() || updated <= 0.0 || updated >= 1.0 {
110                return Err(ScoringError::ArithmeticOverflow(format!(
111                    "gradient update produced invalid base_rate={updated}"
112                )));
113            }
114            Some(updated)
115        } else {
116            None
117        };
118
119        self.alpha = next_alpha;
120        self.beta = next_beta;
121        self.base_rate = next_base_rate;
122        Ok(())
123    }
124
125    pub fn fit(
126        &mut self,
127        raw_scores: &[f64],
128        labels: &[f64],
129        learning_rate: f64,
130        epochs: usize,
131    ) -> ScoringResult<BTreeMap<String, f64>> {
132        if raw_scores.len() != labels.len() {
133            return Err(invalid_input(format!(
134                "raw_scores length {} does not match labels length {}",
135                raw_scores.len(),
136                labels.len()
137            )));
138        }
139        require_finite(learning_rate, "learning_rate")?;
140        if learning_rate <= 0.0 {
141            return Err(invalid_input(format!(
142                "learning_rate must be positive, got {learning_rate}"
143            )));
144        }
145        for (index, raw_score) in raw_scores.iter().copied().enumerate() {
146            require_finite(raw_score, &format!("raw_scores[{index}]"))?;
147        }
148        for (index, label) in labels.iter().copied().enumerate() {
149            require_probability(label, &format!("labels[{index}]"))?;
150        }
151
152        let mut candidate = self.clone();
153        for _ in 0..epochs {
154            for (&raw_score, &label) in raw_scores.iter().zip(labels) {
155                candidate.update(raw_score, label, learning_rate)?;
156            }
157        }
158        *self = candidate;
159        Ok(self.params())
160    }
161
162    pub fn fit_with_options(
163        &mut self,
164        raw_scores: &[f64],
165        labels: &[f64],
166    ) -> ScoringResult<BTreeMap<String, f64>> {
167        self.fit(raw_scores, labels, 0.1, 50)
168    }
169}
170
171impl Default for ParameterLearner {
172    fn default() -> Self {
173        Self {
174            alpha: 1.0,
175            beta: 0.0,
176            base_rate: Some(0.5),
177        }
178    }
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184
185    #[test]
186    fn fit_sharpens_alpha_for_separable_labels() {
187        let mut learner = ParameterLearner::new(0.5, 0.0, None).unwrap();
188        let scores: Vec<f64> = (0_usize..40)
189            .map(|i| if i.is_multiple_of(2) { 5.0 } else { -5.0 })
190            .collect();
191        let labels: Vec<f64> = (0_usize..40)
192            .map(|i| if i.is_multiple_of(2) { 1.0 } else { 0.0 })
193            .collect();
194        let alpha_before = learner.alpha();
195        learner.fit(&scores, &labels, 0.5, 50).unwrap();
196        assert!(
197            learner.alpha() > alpha_before,
198            "alpha {} -> {}",
199            alpha_before,
200            learner.alpha()
201        );
202    }
203
204    #[test]
205    fn update_reduces_logistic_loss() {
206        let mut learner = ParameterLearner::new(1.0, 0.0, Some(0.5)).unwrap();
207        let before = -learner.probability(2.0).unwrap().ln();
208        learner.update(2.0, 1.0, 0.1).unwrap();
209        let after = -learner.probability(2.0).unwrap().ln();
210        assert!(after < before, "loss {before} -> {after}");
211    }
212
213    #[test]
214    fn invalid_updates_leave_parameters_unchanged() {
215        let mut learner = ParameterLearner::default();
216        let before = learner.params();
217        assert!(learner.update(f64::NAN, 1.0, 0.1).is_err());
218        assert_eq!(learner.params(), before);
219
220        assert!(learner.fit(&[1.0, 2.0], &[1.0], 0.1, 1).is_err());
221        assert_eq!(learner.params(), before);
222    }
223}