Skip to main content

zsfm_tabfm/ensemble/
calibration.rs

1//! Platt scaling (binary) and vector scaling (multiclass) output calibration, fit on
2//! out-of-fold probabilities (see `oof.rs`). The wrapper fits these via `scipy.optimize.minimize`
3//! (L-BFGS-B) on a regularized negative-log-likelihood; this ports the same objective but
4//! optimizes it with an in-house box-constrained coordinate descent (golden-section line search
5//! per coordinate) — expect looser numerical agreement than the always-on ensembling path, which
6//! matches the real optimizer's exact trajectory less closely by construction.
7
8const EPS: f64 = 1e-12;
9
10/// Minimizes `f` over `[lo, hi]` via golden-section search (safe under box constraints, unlike
11/// Brent's method's unbounded bracket growth — used here instead of `scalers::brent_minimize`).
12fn golden_section_bounded(f: impl Fn(f64) -> f64, lo: f64, hi: f64) -> f64 {
13    const GR: f64 = 0.618_033_988_749_895; // 1/phi
14    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    /// `p_all`: OOF probabilities `[N][2]`; `y`: true class indices (0 or 1).
47    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    /// `p_all`: OOF probabilities `[N][K]`; `y`: true class indices.
83    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        // Overconfident-but-correct predictions should shrink toward the true labels less
156        // aggressively than raw probs after Platt scaling; here we just check the fit doesn't
157        // diverge and produces valid probabilities.
158        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}