kestrel_chartkit/evaluation/
calibration.rs1use std::collections::HashMap;
7
8use crate::stats::rolling_median;
9
10use super::{SignalEvaluationRecord, TradeOutcome, TradeStats};
11
12#[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 pub brier_score: f64,
27 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
41pub 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#[derive(Debug, Clone, PartialEq)]
104pub struct Cohort {
105 pub key: String,
106 pub stats: TradeStats,
107}
108
109pub 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 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}