zsfm_tabfm/ensemble/
aggregate.rs1pub 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
13pub 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 NnlsWeighted(&'a [f64]),
46 AverageLogits,
49 AverageProbs,
51}
52
53pub 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
74pub fn average_scaled_predictions(scaled: &[f64]) -> f64 {
77 scaled.iter().sum::<f64>() / scaled.len() as f64
78}
79
80pub 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 let n_classes = 3;
101 let shift = 1;
102 let shifted_logits = vec![9.0, 0.1, 0.2];
105 let unshifted = unshift_logits(&shifted_logits, shift, n_classes);
106 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 assert!((via_logits[0] - 0.5).abs() < 1e-6);
117 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}