1pub const PROB_EPSILON: f64 = 1e-10;
14
15#[inline]
16pub fn clamp_prob(p: f64) -> f64 {
17 p.clamp(PROB_EPSILON, 1.0 - PROB_EPSILON)
18}
19
20#[inline]
24pub fn sigmoid(x: f64) -> f64 {
25 if x >= 0.0 {
26 1.0 / (1.0 + (-x).exp())
27 } else {
28 let e = x.exp();
29 e / (1.0 + e)
30 }
31}
32
33#[inline]
36pub fn logit(p: f64) -> f64 {
37 let p = clamp_prob(p);
38 (p / (1.0 - p)).ln()
39}
40
41#[inline]
44pub fn cosine_to_probability(score: f64) -> f64 {
45 clamp_prob(f64::midpoint(1.0, score))
46}
47
48#[inline]
50pub fn prob_not(p: f64) -> f64 {
51 clamp_prob(1.0 - clamp_prob(p))
52}
53
54pub fn prob_and(probs: &[f64]) -> f64 {
56 if probs.is_empty() {
57 return 1.0;
58 }
59 let s: f64 = probs.iter().map(|&p| clamp_prob(p).ln()).sum();
60 s.exp()
61}
62
63pub fn prob_or(probs: &[f64]) -> f64 {
65 if probs.is_empty() {
66 return 0.0;
67 }
68 let s: f64 = probs.iter().map(|&p| (1.0 - clamp_prob(p)).ln()).sum();
69 1.0 - s.exp()
70}
71
72pub fn confidence_scaled_log_odds_pool(probs: &[f64], alpha: f64) -> f64 {
80 if probs.is_empty() {
81 return 0.5;
82 }
83 let n = probs.len() as f64;
84 let mean_logit: f64 = probs.iter().map(|&p| logit(p)).sum::<f64>() / n;
85 sigmoid(mean_logit * n.powf(alpha))
86}
87
88pub fn confidence_scaled_log_odds_pool_weighted(
94 probs: &[f64],
95 weights: &[f64],
96 alpha: f64,
97) -> Result<f64, &'static str> {
98 if probs.len() != weights.len() {
99 return Err("probs and weights must have the same length");
100 }
101 if probs.is_empty() {
102 return Ok(0.5);
103 }
104 if weights.iter().any(|w| *w < 0.0) {
105 return Err("weights must be non-negative");
106 }
107 let sum_w: f64 = weights.iter().sum();
108 if (sum_w - 1.0).abs() > 1e-6 {
109 return Err("weights must sum to 1");
110 }
111 let n = probs.len() as f64;
112 let weighted_logit: f64 = probs.iter().zip(weights).map(|(&p, &w)| w * logit(p)).sum();
113 Ok(sigmoid(n.powf(alpha) * weighted_logit))
114}
115
116#[cfg(test)]
117mod tests {
118 use super::*;
119
120 fn approx_eq(a: f64, b: f64) {
121 assert!((a - b).abs() < 1e-9, "expected {a} ~ {b}");
122 }
123
124 #[test]
125 fn sigmoid_logit_round_trip() {
126 for p in [0.01, 0.1, 0.3, 0.5, 0.7, 0.9, 0.99] {
127 approx_eq(sigmoid(logit(p)), p);
128 }
129 }
130
131 #[test]
132 fn sigmoid_handles_extremes() {
133 assert!(sigmoid(50.0) > 1.0 - 1e-10);
134 assert!(sigmoid(-50.0) < 1e-10);
135 assert!(sigmoid(0.0) - 0.5 < 1e-12);
136 }
137
138 #[test]
139 fn cosine_maps_to_unit_interval() {
140 approx_eq(cosine_to_probability(1.0), 1.0 - PROB_EPSILON);
141 approx_eq(cosine_to_probability(-1.0), PROB_EPSILON);
142 approx_eq(cosine_to_probability(0.0), 0.5);
143 }
144
145 #[test]
146 fn prob_and_log_space_matches_product() {
147 approx_eq(prob_and(&[0.5, 0.5, 0.5]), 0.125);
148 approx_eq(prob_and(&[0.9, 0.8]), 0.72);
149 }
150
151 #[test]
152 fn prob_or_log_space_matches_inclusion_exclusion() {
153 approx_eq(prob_or(&[0.5, 0.5]), 0.75);
154 approx_eq(prob_or(&[0.0, 0.0]), 0.0);
155 }
156
157 #[test]
158 fn confidence_scaled_log_odds_pool_n1_identity() {
159 approx_eq(confidence_scaled_log_odds_pool(&[0.7], 0.5), 0.7);
160 }
161
162 #[test]
163 fn confidence_scaled_log_odds_pool_scale_neutral_at_alpha_zero() {
164 for p in [0.2, 0.5, 0.8] {
168 for n in 1..6 {
169 let probs = vec![p; n];
170 let got = confidence_scaled_log_odds_pool(&probs, 0.0);
171 approx_eq(got, p);
172 }
173 }
174 }
175
176 #[test]
177 fn confidence_scaled_log_odds_pool_amplifies_agreement_at_alpha_half() {
178 let p = 0.7;
183 let p1 = confidence_scaled_log_odds_pool(&[p], 0.5);
184 let p3 = confidence_scaled_log_odds_pool(&[p; 3], 0.5);
185 let p5 = confidence_scaled_log_odds_pool(&[p; 5], 0.5);
186 assert!(p1 < p3 && p3 < p5, "amplification: {p1} < {p3} < {p5}");
187 }
188
189 #[test]
190 fn confidence_scaled_log_odds_pool_irrelevance_preserving() {
191 let probs = [0.2, 0.3, 0.4];
193 let got = confidence_scaled_log_odds_pool(&probs, 0.5);
194 assert!(got < 0.5, "got {got}");
195 }
196
197 #[test]
198 fn confidence_scaled_log_odds_pool_relevance_preserving() {
199 let probs = [0.6, 0.7, 0.8];
200 let got = confidence_scaled_log_odds_pool(&probs, 0.5);
201 assert!(got > 0.5, "got {got}");
202 }
203
204 #[test]
205 fn confidence_scaled_log_odds_pool_symmetric_disagreement_collapses_to_half() {
206 let got = confidence_scaled_log_odds_pool(&[0.3, 0.7], 0.5);
208 approx_eq(got, 0.5);
209 }
210
211 #[test]
212 fn weighted_log_odds_rejects_bad_weights() {
213 assert!(confidence_scaled_log_odds_pool_weighted(&[0.5, 0.5], &[0.5, 0.6], 0.0).is_err());
214 assert!(confidence_scaled_log_odds_pool_weighted(&[0.5, 0.5], &[-0.1, 1.1], 0.0).is_err());
215 }
216}