Skip to main content

zsfm_tabfm/ensemble/
aggregate.rs

1//! Temperature-scaled softmax and the classification/regression ensemble-combination paths from
2//! `TabFMClassifier._process_logits` / `TabFMRegressor._combine_predictions`.
3
4/// `TabFMClassifier.softmax` (temperature divides logits *before* the max-subtracted softmax).
5pub fn softmax_temperature(logits: &[f64], temperature: f64) -> Vec<f64> {
6    let scaled: Vec<f64> = logits.iter().map(|&v| v / temperature).collect();
7    let max = scaled.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
8    let exp: Vec<f64> = scaled.iter().map(|&v| (v - max).exp()).collect();
9    let sum: f64 = exp.iter().sum();
10    exp.into_iter().map(|v| v / sum).collect()
11}
12
13/// Undoes a member's class-label shift: that member was fed `(y_true + shift) % n_classes` as
14/// training labels, so its output logit at position `(c + shift) % n_classes` is the prediction
15/// for original class `c`.
16pub fn unshift_logits(logits: &[f64], shift: usize, n_classes: usize) -> Vec<f64> {
17    (0..n_classes).map(|c| logits[(c + shift) % n_classes]).collect()
18}
19
20fn average_vecs(vecs: &[Vec<f64>]) -> Vec<f64> {
21    let n = vecs.len() as f64;
22    let dim = vecs[0].len();
23    let mut out = vec![0.0; dim];
24    for v in vecs {
25        for (o, x) in out.iter_mut().zip(v.iter()) {
26            *o += x / n;
27        }
28    }
29    out
30}
31
32fn weighted_average_vecs(vecs: &[Vec<f64>], weights: &[f64]) -> Vec<f64> {
33    let dim = vecs[0].len();
34    let mut out = vec![0.0; dim];
35    for (v, &w) in vecs.iter().zip(weights.iter()) {
36        for (o, x) in out.iter_mut().zip(v.iter()) {
37            *o += x * w;
38        }
39    }
40    out
41}
42
43pub enum ClassAggMode<'a> {
44    /// Weighted average of per-member (temperature-softmax'd) probabilities.
45    NnlsWeighted(&'a [f64]),
46    /// Average logits first, then a single temperature-softmax. The wrapper's default
47    /// (`average_logits=True`).
48    AverageLogits,
49    /// Temperature-softmax each member, then plain-average the probabilities.
50    AverageProbs,
51}
52
53/// `_process_logits`: `logits_all` is `[n_estimators][n_classes]`, already un-shifted back to
54/// original class order. Returns final `[n_classes]` probabilities.
55pub fn aggregate_classification(logits_all: &[Vec<f64>], temperature: f64, mode: &ClassAggMode) -> Vec<f64> {
56    match mode {
57        ClassAggMode::NnlsWeighted(weights) => {
58            let probs_all: Vec<Vec<f64>> =
59                logits_all.iter().map(|l| softmax_temperature(l, temperature)).collect();
60            weighted_average_vecs(&probs_all, weights)
61        }
62        ClassAggMode::AverageLogits => {
63            let avg = average_vecs(logits_all);
64            softmax_temperature(&avg, temperature)
65        }
66        ClassAggMode::AverageProbs => {
67            let probs_all: Vec<Vec<f64>> =
68                logits_all.iter().map(|l| softmax_temperature(l, temperature)).collect();
69            average_vecs(&probs_all)
70        }
71    }
72}
73
74/// `_combine_predictions`, NNLS-off (default) path: average the still-*scaled* per-member
75/// predictions first; the caller inverse-transforms the single averaged result afterward.
76pub fn average_scaled_predictions(scaled: &[f64]) -> f64 {
77    scaled.iter().sum::<f64>() / scaled.len() as f64
78}
79
80/// `_combine_predictions`, NNLS-on path: each member's prediction has already been
81/// inverse-transformed by the caller; combine via the fitted weights.
82pub fn weighted_unscaled_predictions(unscaled: &[f64], weights: &[f64]) -> f64 {
83    unscaled.iter().zip(weights.iter()).map(|(p, w)| p * w).sum()
84}
85
86#[cfg(test)]
87mod tests {
88    use super::*;
89
90    #[test]
91    fn test_softmax_sums_to_one() {
92        let p = softmax_temperature(&[1.0, 2.0, 3.0], 0.9);
93        let sum: f64 = p.iter().sum();
94        assert!((sum - 1.0).abs() < 1e-9);
95    }
96
97    #[test]
98    fn test_unshift_roundtrip() {
99        // member fed y'=(y+shift)%3; its logit position for original class c is (c+shift)%3.
100        let n_classes = 3;
101        let shift = 1;
102        // suppose the model's raw (shifted-space) confidence peaks at shifted-position 0,
103        // meaning it's confident about shifted-class 0 = original class (0 - shift) mod 3 = 2.
104        let shifted_logits = vec![9.0, 0.1, 0.2];
105        let unshifted = unshift_logits(&shifted_logits, shift, n_classes);
106        // unshifted[c] = shifted_logits[(c+shift)%3]; unshifted[2] = shifted_logits[(2+1)%3=0] = 9.0
107        assert_eq!(unshifted[2], 9.0);
108    }
109
110    #[test]
111    fn test_average_logits_vs_average_probs_differ() {
112        let logits_all = vec![vec![10.0, 0.0], vec![0.0, 10.0]];
113        let via_logits = aggregate_classification(&logits_all, 0.9, &ClassAggMode::AverageLogits);
114        let via_probs = aggregate_classification(&logits_all, 0.9, &ClassAggMode::AverageProbs);
115        // average-logits of [10,0] and [0,10] -> [5,5] -> softmax -> [0.5,0.5]
116        assert!((via_logits[0] - 0.5).abs() < 1e-6);
117        // average-probs: softmax([10,0]/0.9)~[~1,~0], softmax([0,10]/0.9)~[~0,~1] -> avg ~[0.5,0.5] too here
118        // (symmetric case coincides) — use an asymmetric case to actually differentiate:
119        let logits_all2 = vec![vec![10.0, 0.0], vec![1.0, 0.0]];
120        let l2 = aggregate_classification(&logits_all2, 0.9, &ClassAggMode::AverageLogits);
121        let p2 = aggregate_classification(&logits_all2, 0.9, &ClassAggMode::AverageProbs);
122        assert!((l2[0] - p2[0]).abs() > 1e-3, "expected the two paths to diverge on asymmetric input");
123        let _ = via_probs;
124    }
125}