zsfm_tabfm/ensemble/
calibration.rs1const EPS: f64 = 1e-12;
9
10fn golden_section_bounded(f: impl Fn(f64) -> f64, lo: f64, hi: f64) -> f64 {
13 const GR: f64 = 0.618_033_988_749_895; let (mut a, mut b) = (lo, hi);
15 let mut c = b - GR * (b - a);
16 let mut d = a + GR * (b - a);
17 let mut fc = f(c);
18 let mut fd = f(d);
19 for _ in 0..100 {
20 if (b - a).abs() < 1e-10 {
21 break;
22 }
23 if fc < fd {
24 b = d;
25 d = c;
26 fd = fc;
27 c = b - GR * (b - a);
28 fc = f(c);
29 } else {
30 a = c;
31 c = d;
32 fc = fd;
33 d = a + GR * (b - a);
34 fd = f(d);
35 }
36 }
37 0.5 * (a + b)
38}
39
40pub struct PlattParams {
41 pub a: f64,
42 pub b: f64,
43}
44
45impl PlattParams {
46 pub fn fit(p_all: &[Vec<f64>], y: &[usize], lambda: f64) -> Self {
48 let z: Vec<f64> = p_all.iter().map(|p| ((p[1] + EPS) / (p[0] + EPS)).ln()).collect();
49 let loss = |a: f64, b: f64| -> f64 {
50 let n = z.len() as f64;
51 let mut nll = 0.0;
52 for i in 0..z.len() {
53 let p1 = sigmoid(a * z[i] + b);
54 let p_correct = if y[i] == 1 { p1 } else { 1.0 - p1 };
55 nll -= (p_correct + EPS).ln();
56 }
57 nll / n + lambda * ((a - 1.0).powi(2) + b.powi(2))
58 };
59
60 let mut a = 1.0f64;
61 let mut b = 0.0f64;
62 for _ in 0..20 {
63 a = golden_section_bounded(|av| loss(av, b), 0.8, 1.2);
64 b = golden_section_bounded(|bv| loss(a, bv), -1.0, 1.0);
65 }
66 PlattParams { a, b }
67 }
68
69 pub fn apply(&self, p: &[f64]) -> Vec<f64> {
70 let z = ((p[1] + EPS) / (p[0] + EPS)).ln();
71 let p1 = sigmoid(self.a * z + self.b);
72 vec![1.0 - p1, p1]
73 }
74}
75
76pub struct VectorScalingParams {
77 pub w: Vec<f64>,
78 pub b: Vec<f64>,
79}
80
81impl VectorScalingParams {
82 pub fn fit(p_all: &[Vec<f64>], y: &[usize], lambda: f64) -> Self {
84 let k = p_all[0].len();
85 let z: Vec<Vec<f64>> = p_all.iter().map(|p| p.iter().map(|&v| (v + EPS).ln()).collect()).collect();
86
87 let loss = |w: &[f64], b: &[f64]| -> f64 {
88 let n = z.len() as f64;
89 let mut nll = 0.0;
90 for (i, zi) in z.iter().enumerate() {
91 let logits: Vec<f64> = (0..k).map(|c| w[c] * zi[c] + b[c]).collect();
92 let probs = softmax(&logits);
93 nll -= (probs[y[i]] + EPS).ln();
94 }
95 let reg: f64 =
96 w.iter().map(|&wv| (wv - 1.0).powi(2)).sum::<f64>() + b.iter().map(|&bv| bv.powi(2)).sum::<f64>();
97 nll / n + lambda * reg
98 };
99
100 let mut w = vec![1.0f64; k];
101 let mut b = vec![0.0f64; k];
102 for _ in 0..20 {
103 for c in 0..k {
104 let (w2, b2) = (w.clone(), b.clone());
105 w[c] = golden_section_bounded(
106 |wc| {
107 let mut wt = w2.clone();
108 wt[c] = wc;
109 loss(&wt, &b2)
110 },
111 0.8,
112 1.2,
113 );
114 }
115 for c in 0..k {
116 let (w2, b2) = (w.clone(), b.clone());
117 b[c] = golden_section_bounded(
118 |bc| {
119 let mut bt = b2.clone();
120 bt[c] = bc;
121 loss(&w2, &bt)
122 },
123 -1.0,
124 1.0,
125 );
126 }
127 }
128 VectorScalingParams { w, b }
129 }
130
131 pub fn apply(&self, p: &[f64]) -> Vec<f64> {
132 let k = p.len();
133 let logits: Vec<f64> = (0..k).map(|c| self.w[c] * (p[c] + EPS).ln() + self.b[c]).collect();
134 softmax(&logits)
135 }
136}
137
138fn sigmoid(x: f64) -> f64 {
139 1.0 / (1.0 + (-x).exp())
140}
141
142fn softmax(logits: &[f64]) -> Vec<f64> {
143 let max = logits.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
144 let exp: Vec<f64> = logits.iter().map(|&v| (v - max).exp()).collect();
145 let sum: f64 = exp.iter().sum();
146 exp.into_iter().map(|v| v / sum).collect()
147}
148
149#[cfg(test)]
150mod tests {
151 use super::*;
152
153 #[test]
154 fn test_platt_improves_calibration_of_overconfident_probs() {
155 let p_all = vec![vec![0.99, 0.01], vec![0.02, 0.98], vec![0.6, 0.4], vec![0.4, 0.6]];
159 let y = vec![0, 1, 0, 1];
160 let params = PlattParams::fit(&p_all, &y, 1e-2);
161 for p in &p_all {
162 let out = params.apply(p);
163 assert!((out[0] + out[1] - 1.0).abs() < 1e-6);
164 assert!(out[0] >= 0.0 && out[1] >= 0.0);
165 }
166 }
167
168 #[test]
169 fn test_vector_scaling_valid_probabilities() {
170 let p_all = vec![vec![0.7, 0.2, 0.1], vec![0.1, 0.8, 0.1], vec![0.2, 0.2, 0.6]];
171 let y = vec![0, 1, 2];
172 let params = VectorScalingParams::fit(&p_all, &y, 1e-2);
173 for p in &p_all {
174 let out = params.apply(p);
175 let sum: f64 = out.iter().sum();
176 assert!((sum - 1.0).abs() < 1e-6);
177 }
178 }
179}