Skip to main content

kestrel_chartkit/evaluation/
calibration.rs

1//! Outcome calibration: Brier score, hit-time statistics, agreement-bucketed reliability, and
2//! cohort aggregation over [`super::SignalEvaluationRecord`]s — layered on top of the existing
3//! bar-wise [`super::OutcomeRecorder`]/[`super::TradeStats`], which track realized R-multiples but
4//! not how well a signal's *agreement* predicted its outcome.
5
6use std::collections::HashMap;
7
8use crate::stats::rolling_median;
9
10use super::{SignalEvaluationRecord, TradeOutcome, TradeStats};
11
12/// One agreement bucket's calibration: how predicted agreement compared to the actual win rate
13/// observed within that bucket (a reliability-diagram row).
14#[derive(Debug, Clone, Copy, PartialEq)]
15pub struct CalibrationBucket {
16    pub predicted_mean: f64,
17    pub observed_win_rate: f64,
18    pub count: usize,
19}
20
21#[derive(Debug, Clone, PartialEq)]
22pub struct CalibrationReport {
23    /// Mean squared error between predicted agreement (as a 0.0..=1.0 win probability) and the
24    /// realized binary outcome (`1.0` for `Win`, `0.0` otherwise). Lower is better; `0.0` is
25    /// perfect, `0.25` is what a constant `0.5` forecast scores against a 50/50 outcome mix.
26    pub brier_score: f64,
27    /// Confidence-sorted buckets (a reliability diagram): a well-calibrated signal has
28    /// `observed_win_rate` tracking `predicted_mean` closely in every bucket.
29    pub buckets: Vec<CalibrationBucket>,
30    pub mean_hit_time_bars: f64,
31    pub median_hit_time_bars: f64,
32}
33
34fn outcome_as_probability(outcome: TradeOutcome) -> f64 {
35    match outcome {
36        TradeOutcome::Win => 1.0,
37        TradeOutcome::Loss | TradeOutcome::BreakEven | TradeOutcome::Expired => 0.0,
38    }
39}
40
41/// Computes a [`CalibrationReport`] from `records`, using `agreement` (expected `0.0..=1.0`) as
42/// each record's predicted win probability. `num_buckets` controls the reliability-diagram
43/// resolution (records are sorted by agreement and split into roughly equal-sized buckets).
44pub fn compute_calibration(
45    records: &[SignalEvaluationRecord],
46    num_buckets: usize,
47) -> CalibrationReport {
48    if records.is_empty() {
49        return CalibrationReport {
50            brier_score: 0.0,
51            buckets: Vec::new(),
52            mean_hit_time_bars: 0.0,
53            median_hit_time_bars: 0.0,
54        };
55    }
56
57    let brier_score = records
58        .iter()
59        .map(|r| {
60            let p = r.agreement.clamp(0.0, 1.0);
61            let o = outcome_as_probability(r.outcome);
62            (p - o).powi(2)
63        })
64        .sum::<f64>()
65        / records.len() as f64;
66
67    let mut sorted: Vec<&SignalEvaluationRecord> = records.iter().collect();
68    sorted.sort_by(|a, b| a.agreement.total_cmp(&b.agreement));
69
70    let num_buckets = num_buckets.max(1).min(sorted.len());
71    let bucket_size = sorted.len().div_ceil(num_buckets);
72    let buckets = sorted
73        .chunks(bucket_size.max(1))
74        .map(|chunk| {
75            let predicted_mean =
76                chunk.iter().map(|r| r.agreement).sum::<f64>() / chunk.len() as f64;
77            let wins = chunk
78                .iter()
79                .filter(|r| r.outcome == TradeOutcome::Win)
80                .count();
81            CalibrationBucket {
82                predicted_mean,
83                observed_win_rate: wins as f64 / chunk.len() as f64,
84                count: chunk.len(),
85            }
86        })
87        .collect();
88
89    let durations: Vec<f64> = records.iter().map(|r| r.duration_bars as f64).collect();
90    let mean_hit_time_bars = durations.iter().sum::<f64>() / durations.len() as f64;
91    let median_hit_time_bars = rolling_median(&durations);
92
93    CalibrationReport {
94        brier_score,
95        buckets,
96        mean_hit_time_bars,
97        median_hit_time_bars,
98    }
99}
100
101/// One cohort's aggregated [`TradeStats`], keyed by a caller-supplied grouping label (e.g. trade
102/// direction, session, or any other categorical dimension).
103#[derive(Debug, Clone, PartialEq)]
104pub struct Cohort {
105    pub key: String,
106    pub stats: TradeStats,
107}
108
109/// Groups `records` by `key_fn` and computes [`TradeStats`] per cohort, sorted by key for
110/// deterministic output.
111pub fn cohort_aggregate<K, F>(records: &[SignalEvaluationRecord], key_fn: F) -> Vec<Cohort>
112where
113    K: ToString,
114    F: Fn(&SignalEvaluationRecord) -> K,
115{
116    let mut groups: HashMap<String, Vec<SignalEvaluationRecord>> = HashMap::new();
117    for record in records {
118        groups
119            .entry(key_fn(record).to_string())
120            .or_default()
121            .push(record.clone());
122    }
123
124    let mut cohorts: Vec<Cohort> = groups
125        .into_iter()
126        .map(|(key, group_records)| Cohort {
127            key,
128            stats: TradeStats::compute(&group_records),
129        })
130        .collect();
131    cohorts.sort_by(|a, b| a.key.cmp(&b.key));
132    cohorts
133}
134
135#[cfg(test)]
136mod tests {
137    use super::*;
138    use crate::signal::TriggerAction;
139
140    fn record(agreement: f64, outcome: TradeOutcome, duration_bars: u32) -> SignalEvaluationRecord {
141        SignalEvaluationRecord {
142            timestamp: 0,
143            trigger: TriggerAction::Buy,
144            score: 0.5,
145            agreement,
146            entry_price: 100.0,
147            exit_price: 101.0,
148            realized_r_multiple: 1.0,
149            duration_bars,
150            outcome,
151        }
152    }
153
154    #[test]
155    fn test_brier_score_zero_for_perfect_forecasts() {
156        let records = vec![
157            record(1.0, TradeOutcome::Win, 5),
158            record(0.0, TradeOutcome::Loss, 5),
159        ];
160        let report = compute_calibration(&records, 2);
161        assert!(report.brier_score < 1e-9);
162    }
163
164    #[test]
165    fn test_brier_score_positive_for_overconfident_forecasts() {
166        let records = vec![
167            record(0.9, TradeOutcome::Loss, 5),
168            record(0.9, TradeOutcome::Loss, 5),
169        ];
170        let report = compute_calibration(&records, 2);
171        assert!((report.brier_score - 0.81).abs() < 1e-9);
172    }
173
174    #[test]
175    fn test_buckets_reveal_miscalibration() {
176        // High agreement but only ever loses: the bucket's observed win rate must diverge
177        // sharply from its predicted mean.
178        let records: Vec<_> = (0..10)
179            .map(|_| record(0.9, TradeOutcome::Loss, 3))
180            .collect();
181        let report = compute_calibration(&records, 1);
182        assert_eq!(report.buckets.len(), 1);
183        let bucket = report.buckets[0];
184        assert!((bucket.predicted_mean - 0.9).abs() < 1e-9);
185        assert_eq!(bucket.observed_win_rate, 0.0);
186    }
187
188    #[test]
189    fn test_hit_time_statistics() {
190        let records = vec![
191            record(0.5, TradeOutcome::Win, 2),
192            record(0.5, TradeOutcome::Win, 4),
193            record(0.5, TradeOutcome::Loss, 6),
194        ];
195        let report = compute_calibration(&records, 1);
196        assert!((report.mean_hit_time_bars - 4.0).abs() < 1e-9);
197        assert!((report.median_hit_time_bars - 4.0).abs() < 1e-9);
198    }
199
200    #[test]
201    fn test_cohort_aggregate_groups_and_sorts_by_key() {
202        let records = vec![
203            record(0.6, TradeOutcome::Win, 3),
204            record(0.4, TradeOutcome::Loss, 5),
205            record(0.7, TradeOutcome::Win, 2),
206        ];
207        let cohorts = cohort_aggregate(
208            &records,
209            |r| if r.agreement >= 0.5 { "high" } else { "low" },
210        );
211
212        assert_eq!(cohorts.len(), 2);
213        assert_eq!(cohorts[0].key, "high");
214        assert_eq!(cohorts[0].stats.total_trades, 2);
215        assert_eq!(cohorts[1].key, "low");
216        assert_eq!(cohorts[1].stats.total_trades, 1);
217    }
218
219    #[test]
220    fn test_empty_records_produce_neutral_report() {
221        let report = compute_calibration(&[], 5);
222        assert_eq!(report.brier_score, 0.0);
223        assert!(report.buckets.is_empty());
224    }
225}