uqa_scoring/
parameter_learner.rs1use 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 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}