use std::collections::HashMap;
use crate::stats::rolling_median;
use super::{SignalEvaluationRecord, TradeOutcome, TradeStats};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CalibrationBucket {
pub predicted_mean: f64,
pub observed_win_rate: f64,
pub count: usize,
}
#[derive(Debug, Clone, PartialEq)]
pub struct CalibrationReport {
pub brier_score: f64,
pub buckets: Vec<CalibrationBucket>,
pub mean_hit_time_bars: f64,
pub median_hit_time_bars: f64,
}
fn outcome_as_probability(outcome: TradeOutcome) -> f64 {
match outcome {
TradeOutcome::Win => 1.0,
TradeOutcome::Loss | TradeOutcome::BreakEven | TradeOutcome::Expired => 0.0,
}
}
pub fn compute_calibration(
records: &[SignalEvaluationRecord],
num_buckets: usize,
) -> CalibrationReport {
if records.is_empty() {
return CalibrationReport {
brier_score: 0.0,
buckets: Vec::new(),
mean_hit_time_bars: 0.0,
median_hit_time_bars: 0.0,
};
}
let brier_score = records
.iter()
.map(|r| {
let p = r.agreement.clamp(0.0, 1.0);
let o = outcome_as_probability(r.outcome);
(p - o).powi(2)
})
.sum::<f64>()
/ records.len() as f64;
let mut sorted: Vec<&SignalEvaluationRecord> = records.iter().collect();
sorted.sort_by(|a, b| a.agreement.total_cmp(&b.agreement));
let num_buckets = num_buckets.max(1).min(sorted.len());
let bucket_size = sorted.len().div_ceil(num_buckets);
let buckets = sorted
.chunks(bucket_size.max(1))
.map(|chunk| {
let predicted_mean =
chunk.iter().map(|r| r.agreement).sum::<f64>() / chunk.len() as f64;
let wins = chunk
.iter()
.filter(|r| r.outcome == TradeOutcome::Win)
.count();
CalibrationBucket {
predicted_mean,
observed_win_rate: wins as f64 / chunk.len() as f64,
count: chunk.len(),
}
})
.collect();
let durations: Vec<f64> = records.iter().map(|r| r.duration_bars as f64).collect();
let mean_hit_time_bars = durations.iter().sum::<f64>() / durations.len() as f64;
let median_hit_time_bars = rolling_median(&durations);
CalibrationReport {
brier_score,
buckets,
mean_hit_time_bars,
median_hit_time_bars,
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Cohort {
pub key: String,
pub stats: TradeStats,
}
pub fn cohort_aggregate<K, F>(records: &[SignalEvaluationRecord], key_fn: F) -> Vec<Cohort>
where
K: ToString,
F: Fn(&SignalEvaluationRecord) -> K,
{
let mut groups: HashMap<String, Vec<SignalEvaluationRecord>> = HashMap::new();
for record in records {
groups
.entry(key_fn(record).to_string())
.or_default()
.push(record.clone());
}
let mut cohorts: Vec<Cohort> = groups
.into_iter()
.map(|(key, group_records)| Cohort {
key,
stats: TradeStats::compute(&group_records),
})
.collect();
cohorts.sort_by(|a, b| a.key.cmp(&b.key));
cohorts
}
#[cfg(test)]
mod tests {
use super::*;
use crate::signal::TriggerAction;
fn record(agreement: f64, outcome: TradeOutcome, duration_bars: u32) -> SignalEvaluationRecord {
SignalEvaluationRecord {
timestamp: 0,
trigger: TriggerAction::Buy,
score: 0.5,
agreement,
entry_price: 100.0,
exit_price: 101.0,
realized_r_multiple: 1.0,
duration_bars,
outcome,
}
}
#[test]
fn test_brier_score_zero_for_perfect_forecasts() {
let records = vec![
record(1.0, TradeOutcome::Win, 5),
record(0.0, TradeOutcome::Loss, 5),
];
let report = compute_calibration(&records, 2);
assert!(report.brier_score < 1e-9);
}
#[test]
fn test_brier_score_positive_for_overconfident_forecasts() {
let records = vec![
record(0.9, TradeOutcome::Loss, 5),
record(0.9, TradeOutcome::Loss, 5),
];
let report = compute_calibration(&records, 2);
assert!((report.brier_score - 0.81).abs() < 1e-9);
}
#[test]
fn test_buckets_reveal_miscalibration() {
let records: Vec<_> = (0..10)
.map(|_| record(0.9, TradeOutcome::Loss, 3))
.collect();
let report = compute_calibration(&records, 1);
assert_eq!(report.buckets.len(), 1);
let bucket = report.buckets[0];
assert!((bucket.predicted_mean - 0.9).abs() < 1e-9);
assert_eq!(bucket.observed_win_rate, 0.0);
}
#[test]
fn test_hit_time_statistics() {
let records = vec![
record(0.5, TradeOutcome::Win, 2),
record(0.5, TradeOutcome::Win, 4),
record(0.5, TradeOutcome::Loss, 6),
];
let report = compute_calibration(&records, 1);
assert!((report.mean_hit_time_bars - 4.0).abs() < 1e-9);
assert!((report.median_hit_time_bars - 4.0).abs() < 1e-9);
}
#[test]
fn test_cohort_aggregate_groups_and_sorts_by_key() {
let records = vec![
record(0.6, TradeOutcome::Win, 3),
record(0.4, TradeOutcome::Loss, 5),
record(0.7, TradeOutcome::Win, 2),
];
let cohorts = cohort_aggregate(
&records,
|r| if r.agreement >= 0.5 { "high" } else { "low" },
);
assert_eq!(cohorts.len(), 2);
assert_eq!(cohorts[0].key, "high");
assert_eq!(cohorts[0].stats.total_trades, 2);
assert_eq!(cohorts[1].key, "low");
assert_eq!(cohorts[1].stats.total_trades, 1);
}
#[test]
fn test_empty_records_produce_neutral_report() {
let report = compute_calibration(&[], 5);
assert_eq!(report.brier_score, 0.0);
assert!(report.buckets.is_empty());
}
}