use std::fmt;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct CalibratedProbability {
pub raw_score: f64,
pub probability: f64,
pub model_version: String,
pub train_sample_size: usize,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct IsotonicCalibrator {
pub thresholds: Vec<f64>,
pub probabilities: Vec<f64>,
pub train_sample_size: usize,
}
impl IsotonicCalibrator {
pub fn fit(samples: &[(f64, bool)]) -> Option<Self> {
if samples.is_empty() {
return None;
}
let mut sorted = samples.to_vec();
sorted.sort_by(|a, b| a.0.total_cmp(&b.0));
let mut scores: Vec<f64> = Vec::with_capacity(sorted.len());
let mut weights: Vec<f64> = Vec::with_capacity(sorted.len());
let mut values: Vec<f64> = Vec::with_capacity(sorted.len());
for (score, outcome) in sorted {
scores.push(score);
weights.push(1.0);
values.push(if outcome { 1.0 } else { 0.0 });
}
let mut i = 0;
while i < values.len() {
if i > 0 && values[i] < values[i - 1] {
let w1 = weights[i - 1];
let w2 = weights[i];
let v1 = values[i - 1];
let v2 = values[i];
let w_new = w1 + w2;
let v_new = (w1 * v1 + w2 * v2) / w_new;
weights[i - 1] = w_new;
values[i - 1] = v_new;
scores[i - 1] = scores[i];
weights.remove(i);
values.remove(i);
scores.remove(i);
i -= 1; } else {
i += 1;
}
}
Some(Self {
thresholds: scores,
probabilities: values,
train_sample_size: samples.len(),
})
}
pub fn predict(&self, raw_score: f64) -> f64 {
if self.thresholds.is_empty() {
return 0.5;
}
if raw_score <= self.thresholds[0] {
return self.probabilities[0].clamp(0.0, 1.0);
}
let last_idx = self.thresholds.len() - 1;
if raw_score >= self.thresholds[last_idx] {
return self.probabilities[last_idx].clamp(0.0, 1.0);
}
match self
.thresholds
.binary_search_by(|t| t.total_cmp(&raw_score))
{
Ok(idx) => self.probabilities[idx].clamp(0.0, 1.0),
Err(idx) => {
let x0 = self.thresholds[idx - 1];
let x1 = self.thresholds[idx];
let y0 = self.probabilities[idx - 1];
let y1 = self.probabilities[idx];
let frac = if (x1 - x0).abs() > 1e-12 {
(raw_score - x0) / (x1 - x0)
} else {
0.0
};
(y0 + frac * (y1 - y0)).clamp(0.0, 1.0)
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct CalibrationMetrics {
pub brier_score: f64,
pub baseline_brier_score: f64,
pub brier_skill_score: f64,
pub log_loss: f64,
pub expected_calibration_error: f64,
pub sample_size: usize,
}
impl fmt::Display for CalibrationMetrics {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"Brier: {:.4} (BSS: {:.2}%), LogLoss: {:.4}, ECE: {:.4} (N={})",
self.brier_score,
self.brier_skill_score * 100.0,
self.log_loss,
self.expected_calibration_error,
self.sample_size
)
}
}
pub fn compute_calibration_metrics(
predicted: &[f64],
actual: &[bool],
baseline_rate: f64,
num_bins: usize,
) -> Option<CalibrationMetrics> {
if predicted.len() != actual.len() || predicted.is_empty() {
return None;
}
let n = predicted.len() as f64;
let baseline = baseline_rate.clamp(1e-6, 1.0 - 1e-6);
let mut brier_sum = 0.0;
let mut baseline_brier_sum = 0.0;
let mut log_loss_sum = 0.0;
for (&p_raw, &y) in predicted.iter().zip(actual.iter()) {
let p = p_raw.clamp(1e-15, 1.0 - 1e-15);
let target = if y { 1.0 } else { 0.0 };
brier_sum += (p - target).powi(2);
baseline_brier_sum += (baseline - target).powi(2);
let log_p = if y { p.ln() } else { (1.0 - p).ln() };
log_loss_sum -= log_p;
}
let brier_score = brier_sum / n;
let baseline_brier_score = baseline_brier_sum / n;
let brier_skill_score = if baseline_brier_score > 1e-12 {
1.0 - (brier_score / baseline_brier_score)
} else {
0.0
};
let log_loss = log_loss_sum / n;
let bins = num_bins.max(1);
let mut bin_counts = vec![0usize; bins];
let mut bin_pred_sums = vec![0.0f64; bins];
let mut bin_actual_sums = vec![0.0f64; bins];
for (&p_raw, &y) in predicted.iter().zip(actual.iter()) {
let p = p_raw.clamp(0.0, 1.0);
let bin_idx = ((p * bins as f64).floor() as usize).min(bins - 1);
bin_counts[bin_idx] += 1;
bin_pred_sums[bin_idx] += p;
bin_actual_sums[bin_idx] += if y { 1.0 } else { 0.0 };
}
let mut ece = 0.0;
for i in 0..bins {
if bin_counts[i] > 0 {
let count = bin_counts[i] as f64;
let avg_pred = bin_pred_sums[i] / count;
let avg_actual = bin_actual_sums[i] / count;
ece += (count / n) * (avg_pred - avg_actual).abs();
}
}
Some(CalibrationMetrics {
brier_score,
baseline_brier_score,
brier_skill_score,
log_loss,
expected_calibration_error: ece,
sample_size: predicted.len(),
})
}
pub fn block_bootstrap_brier(
predicted: &[f64],
actual: &[bool],
block_size: usize,
num_bootstraps: usize,
seed: u64,
) -> Option<(f64, f64, f64)> {
let n = predicted.len();
if n == 0 || block_size == 0 || num_bootstraps == 0 {
return None;
}
let mut lcg = seed.wrapping_add(1);
let mut bootstrap_scores = Vec::with_capacity(num_bootstraps);
for _ in 0..num_bootstraps {
let mut sample_brier_sum = 0.0;
let mut count = 0;
while count < n {
lcg = lcg.wrapping_mul(6364136223846793005).wrapping_add(1);
let start = (lcg as usize) % n;
let take = block_size.min(n - count);
for k in 0..take {
let idx = (start + k) % n;
let p = predicted[idx].clamp(0.0, 1.0);
let y = if actual[idx] { 1.0 } else { 0.0 };
sample_brier_sum += (p - y).powi(2);
}
count += take;
}
bootstrap_scores.push(sample_brier_sum / count as f64);
}
bootstrap_scores.sort_by(|a, b| a.total_cmp(b));
let mean = bootstrap_scores.iter().sum::<f64>() / num_bootstraps as f64;
let p05_idx = ((0.05 * num_bootstraps as f64).floor() as usize).min(num_bootstraps - 1);
let p95_idx = ((0.95 * num_bootstraps as f64).floor() as usize).min(num_bootstraps - 1);
Some((mean, bootstrap_scores[p05_idx], bootstrap_scores[p95_idx]))
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct ValidationExperimentManifest {
pub experiment_id: String,
pub model_version: String,
pub train_range: (usize, usize),
pub test_range: (usize, usize),
pub embargo_bars: usize,
pub train_sample_size: usize,
pub test_sample_size: usize,
pub metrics: CalibrationMetrics,
}