1use anyhow::{Context, Result};
10use serde::{Deserialize, Serialize};
11use std::collections::HashMap;
12use std::sync::Arc;
13use tokio::sync::RwLock;
14use tracing::{debug, info, warn};
15
16pub struct CalibrationSystem {
18 config: CalibrationConfig,
19 baseline_calibration: Arc<RwLock<Option<CalibrationMeasurement>>>,
21 current_calibration: Arc<RwLock<CalibrationMeasurement>>,
23 temperature_models: Arc<RwLock<HashMap<String, TemperatureModel>>>,
25 feature_scalers: Arc<RwLock<FeatureScalerSet>>,
27 calibration_history: Arc<RwLock<Vec<CalibrationSnapshot>>>,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize)]
32pub struct CalibrationConfig {
33 pub max_ece_drift: f32,
35 pub log_odds_cap: f32,
37 pub temperature: f32,
39 pub min_samples_for_calibration: usize,
41 pub measurement_window_size: usize,
43 pub auto_temperature_adjustment: bool,
45}
46
47#[derive(Debug, Clone, Serialize, Deserialize)]
49pub struct CalibrationMeasurement {
50 pub ece: f32,
52 pub ece_ci_lower: f32,
54 pub ece_ci_upper: f32,
55 pub mce: f32,
57 pub reliability_data: Vec<ReliabilityPoint>,
59 pub sample_count: usize,
61 pub timestamp: std::time::SystemTime,
63 pub by_query_type: HashMap<String, f32>,
65 pub by_language: HashMap<String, f32>,
67}
68
69#[derive(Debug, Clone, Serialize, Deserialize)]
71pub struct ReliabilityPoint {
72 pub predicted_prob: f32,
74 pub actual_positive_rate: f32,
76 pub bin_count: usize,
78 pub confidence_interval: (f32, f32),
80}
81
82#[derive(Debug, Clone, Serialize, Deserialize)]
84pub struct TemperatureModel {
85 pub temperature: f32,
87 pub bias: f32,
89 pub ece_before: f32,
91 pub ece_after: f32,
92 pub training_samples: usize,
94 pub scope: String,
96}
97
98#[derive(Debug, Clone, Serialize, Deserialize)]
100pub struct FeatureScalerSet {
101 pub semantic_feature_caps: HashMap<String, f32>,
103 pub linear_scalers: HashMap<String, LinearScaler>,
105 pub robust_scalers: HashMap<String, RobustScaler>,
107}
108
109#[derive(Debug, Clone, Serialize, Deserialize)]
110pub struct LinearScaler {
111 pub mean: f32,
112 pub std: f32,
113}
114
115#[derive(Debug, Clone, Serialize, Deserialize)]
116pub struct RobustScaler {
117 pub median: f32,
118 pub iqr: f32, }
120
121#[derive(Debug, Clone, Serialize, Deserialize)]
123pub struct CalibrationSnapshot {
124 pub timestamp: std::time::SystemTime,
125 pub ece: f32,
126 pub mce: f32,
127 pub sample_count: usize,
128 pub semantic_features_active: bool,
129 pub temperature: f32,
130}
131
132#[derive(Debug, Clone)]
134pub struct CalibratedPrediction {
135 pub raw_score: f32,
136 pub calibrated_score: f32,
137 pub confidence: f32,
138 pub query_type: String,
139 pub language: Option<String>,
140 pub features_used: Vec<String>,
141}
142
143#[derive(Debug, Clone)]
145pub struct CalibrationSample {
146 pub prediction: f32,
147 pub actual_relevance: f32, pub query_type: String,
149 pub language: Option<String>,
150}
151
152impl CalibrationSystem {
153 pub async fn new(config: CalibrationConfig) -> Result<Self> {
155 info!("Creating calibration preservation system");
156 info!("ECE drift limit: {:.4}, log-odds cap: {:.2}, temperature: {:.2}",
157 config.max_ece_drift, config.log_odds_cap, config.temperature);
158
159 Ok(Self {
160 config,
161 baseline_calibration: Arc::new(RwLock::new(None)),
162 current_calibration: Arc::new(RwLock::new(CalibrationMeasurement::default())),
163 temperature_models: Arc::new(RwLock::new(HashMap::new())),
164 feature_scalers: Arc::new(RwLock::new(FeatureScalerSet::default())),
165 calibration_history: Arc::new(RwLock::new(Vec::new())),
166 })
167 }
168
169 pub async fn establish_baseline(&self, samples: &[CalibrationSample]) -> Result<()> {
171 info!("Establishing baseline calibration with {} samples", samples.len());
172
173 if samples.len() < self.config.min_samples_for_calibration {
174 anyhow::bail!("Insufficient samples for baseline: {} < {}",
175 samples.len(), self.config.min_samples_for_calibration);
176 }
177
178 let measurement = self.calculate_calibration_measurement(samples).await?;
180
181 info!("Baseline calibration established: ECE = {:.4}, MCE = {:.4}",
182 measurement.ece, measurement.mce);
183
184 *self.baseline_calibration.write().await = Some(measurement.clone());
185 self.add_calibration_snapshot(measurement.ece, measurement.mce, false, self.config.temperature).await;
187
188 *self.current_calibration.write().await = measurement;
190
191 Ok(())
192 }
193
194 pub async fn scale_features(&self, features: &HashMap<String, f32>) -> Result<HashMap<String, f32>> {
196 let scalers = self.feature_scalers.read().await;
197 let mut scaled_features = HashMap::new();
198
199 for (name, value) in features {
200 let scaled_value = if name.contains("semantic") {
201 let capped_value = if let Some(&cap) = scalers.semantic_feature_caps.get(name) {
203 value.clamp(-cap, cap)
204 } else {
205 value.clamp(-self.config.log_odds_cap, self.config.log_odds_cap)
206 };
207
208 if let Some(robust_scaler) = scalers.robust_scalers.get(name) {
210 (capped_value - robust_scaler.median) / (robust_scaler.iqr + 1e-8)
211 } else {
212 capped_value
213 }
214 } else {
215 if let Some(linear_scaler) = scalers.linear_scalers.get(name) {
217 (value - linear_scaler.mean) / (linear_scaler.std + 1e-8)
218 } else {
219 *value
220 }
221 };
222
223 scaled_features.insert(name.clone(), scaled_value);
224 }
225
226 debug!("Scaled {} features with calibration preservation", scaled_features.len());
227 Ok(scaled_features)
228 }
229
230 pub async fn apply_temperature_scaling(&self, prediction: f32, query_type: &str, language: Option<&str>) -> Result<f32> {
232 let models = self.temperature_models.read().await;
233
234 let scope_key = match language {
236 Some(lang) => format!("{}:{}", query_type, lang),
237 None => query_type.to_string(),
238 };
239
240 let temperature = if let Some(model) = models.get(&scope_key) {
241 model.temperature
242 } else if let Some(model) = models.get(query_type) {
243 model.temperature
244 } else {
245 self.config.temperature
246 };
247
248 let logit_p = (prediction / (1.0 - prediction + 1e-8)).ln();
250 let scaled_logit = logit_p / temperature;
251 let calibrated_p = 1.0 / (1.0 + (-scaled_logit).exp());
252
253 Ok(calibrated_p.clamp(0.001, 0.999)) }
255
256 pub async fn update_calibration(&self, samples: &[CalibrationSample]) -> Result<CalibrationStatus> {
258 if samples.len() < self.config.min_samples_for_calibration / 4 {
259 return Ok(CalibrationStatus::InsufficientData);
260 }
261
262 let current_measurement = self.calculate_calibration_measurement(samples).await?;
264
265 let baseline_guard = self.baseline_calibration.read().await;
267 let drift_status = if let Some(baseline) = baseline_guard.as_ref() {
268 let ece_drift = (current_measurement.ece - baseline.ece).abs();
269
270 if ece_drift > self.config.max_ece_drift {
271 warn!("Calibration drift detected: ECE drift {:.4} > limit {:.4}",
272 ece_drift, self.config.max_ece_drift);
273 CalibrationStatus::DriftDetected {
274 ece_drift,
275 current_ece: current_measurement.ece,
276 baseline_ece: baseline.ece,
277 }
278 } else {
279 CalibrationStatus::WithinLimits {
280 ece_drift,
281 current_ece: current_measurement.ece,
282 }
283 }
284 } else {
285 CalibrationStatus::NoBaseline
286 };
287
288 *self.current_calibration.write().await = current_measurement.clone();
290
291 self.add_calibration_snapshot(current_measurement.ece, current_measurement.mce, true, self.config.temperature).await;
293
294 if matches!(drift_status, CalibrationStatus::DriftDetected { .. }) && self.config.auto_temperature_adjustment {
296 self.adjust_temperature_models(samples).await?;
297 }
298
299 info!("Calibration updated: ECE = {:.4}, status = {:?}",
300 current_measurement.ece, drift_status);
301
302 Ok(drift_status)
303 }
304
305 pub async fn train_temperature_models(&self, samples: &[CalibrationSample]) -> Result<()> {
307 info!("Training temperature scaling models with {} samples", samples.len());
308
309 let mut grouped_samples: HashMap<String, Vec<&CalibrationSample>> = HashMap::new();
311
312 for sample in samples {
313 let key = match &sample.language {
314 Some(lang) => format!("{}:{}", sample.query_type, lang),
315 None => sample.query_type.clone(),
316 };
317 grouped_samples.entry(key.clone()).or_default().push(sample);
318
319 grouped_samples.entry(sample.query_type.clone()).or_default().push(sample);
321 }
322
323 let mut models = self.temperature_models.write().await;
324
325 for (scope, group_samples) in grouped_samples {
326 if group_samples.len() >= self.config.min_samples_for_calibration / 4 {
327 let model = self.fit_temperature_model(group_samples, &scope).await?;
328 models.insert(scope, model);
329 }
330 }
331
332 info!("Trained {} temperature models", models.len());
333 Ok(())
334 }
335
336 pub async fn get_calibration_metrics(&self) -> CalibrationMetrics {
338 let current = self.current_calibration.read().await;
339 let baseline = self.baseline_calibration.read().await;
340 let history = self.calibration_history.read().await;
341
342 let ece_drift = if let Some(baseline_measurement) = baseline.as_ref() {
343 (current.ece - baseline_measurement.ece).abs()
344 } else {
345 0.0
346 };
347
348 let recent_trend = if history.len() >= 5 {
349 let recent_eces: Vec<f32> = history.iter().rev().take(5).map(|s| s.ece).collect();
350 let first = recent_eces[4];
351 let last = recent_eces[0];
352 (last - first) / 5.0 } else {
354 0.0
355 };
356
357 CalibrationMetrics {
358 current_ece: current.ece,
359 baseline_ece: baseline.as_ref().map(|b| b.ece).unwrap_or(0.0),
360 ece_drift,
361 max_allowed_drift: self.config.max_ece_drift,
362 within_limits: ece_drift <= self.config.max_ece_drift,
363 recent_trend,
364 measurement_count: history.len(),
365 sample_count: current.sample_count,
366 }
367 }
368
369 fn calculate_calibration_measurement<'a>(&'a self, samples: &'a [CalibrationSample]) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<CalibrationMeasurement>> + Send + 'a>> {
372 Box::pin(async move {
373 let num_bins = 10;
374 let mut bins = vec![Vec::new(); num_bins];
375
376 for sample in samples {
378 let bin_idx = ((sample.prediction * num_bins as f32) as usize).min(num_bins - 1);
379 bins[bin_idx].push(sample);
380 }
381
382 let mut reliability_data = Vec::new();
384 let mut total_ece = 0.0;
385 let mut total_samples = 0;
386 let mut max_calibration_error = 0.0;
387
388 for (i, bin) in bins.iter().enumerate() {
389 if bin.is_empty() {
390 continue;
391 }
392
393 let bin_center = (i as f32 + 0.5) / num_bins as f32;
394 let actual_positive_rate = bin.iter()
395 .map(|s| s.actual_relevance)
396 .sum::<f32>() / bin.len() as f32;
397
398 let bin_error = (bin_center - actual_positive_rate).abs();
399 let bin_weight = bin.len() as f32 / samples.len() as f32;
400
401 total_ece += bin_error * bin_weight;
402 max_calibration_error = f32::max(max_calibration_error, bin_error);
403 total_samples += bin.len();
404
405 let (ci_lower, ci_upper) = self.calculate_wilson_confidence_interval(
407 actual_positive_rate, bin.len()
408 );
409
410 reliability_data.push(ReliabilityPoint {
411 predicted_prob: bin_center,
412 actual_positive_rate,
413 bin_count: bin.len(),
414 confidence_interval: (ci_lower, ci_upper),
415 });
416 }
417
418 let ece_std_error = (total_ece * (1.0 - total_ece) / samples.len() as f32).sqrt();
420 let ece_ci_lower = (total_ece - 1.96 * ece_std_error).max(0.0);
421 let ece_ci_upper = (total_ece + 1.96 * ece_std_error).min(1.0);
422
423 let by_query_type = self.calculate_ece_by_category(samples, |s| &s.query_type).await?;
425 let by_language = self.calculate_ece_by_language(samples).await?;
426
427 Ok(CalibrationMeasurement {
428 ece: total_ece,
429 ece_ci_lower,
430 ece_ci_upper,
431 mce: max_calibration_error,
432 reliability_data,
433 sample_count: samples.len(),
434 timestamp: std::time::SystemTime::now(),
435 by_query_type,
436 by_language,
437 })
438 })
439 }
440
441 fn calculate_wilson_confidence_interval(&self, p: f32, n: usize) -> (f32, f32) {
442 if n == 0 {
443 return (0.0, 1.0);
444 }
445
446 let z = 1.96; let n_f = n as f32;
448
449 let center = p + z * z / (2.0 * n_f);
450 let half_width = z * ((p * (1.0 - p) / n_f) + (z * z / (4.0 * n_f * n_f))).sqrt();
451 let denominator = 1.0 + z * z / n_f;
452
453 let lower = ((center - half_width) / denominator).max(0.0);
454 let upper = ((center + half_width) / denominator).min(1.0);
455
456 (lower, upper)
457 }
458
459 async fn calculate_ece_by_category<F>(&self, samples: &[CalibrationSample], category_fn: F) -> Result<HashMap<String, f32>>
460 where
461 F: Fn(&CalibrationSample) -> &String,
462 {
463 let mut category_groups: HashMap<String, Vec<&CalibrationSample>> = HashMap::new();
464
465 for sample in samples {
466 let category = category_fn(sample);
467 category_groups.entry(category.clone()).or_default().push(sample);
468 }
469
470 let mut category_eces = HashMap::new();
471
472 for (category, group) in category_groups {
473 if group.len() >= 10 { let group_samples: Vec<CalibrationSample> = group.into_iter().cloned().collect();
475 let measurement = self.calculate_calibration_measurement(&group_samples).await?;
476 category_eces.insert(category, measurement.ece);
477 }
478 }
479
480 Ok(category_eces)
481 }
482
483 async fn calculate_ece_by_language(&self, samples: &[CalibrationSample]) -> Result<HashMap<String, f32>> {
484 let mut language_groups: HashMap<String, Vec<&CalibrationSample>> = HashMap::new();
485
486 for sample in samples {
487 let language = sample.language.as_deref().unwrap_or("unknown");
488 language_groups.entry(language.to_string()).or_default().push(sample);
489 }
490
491 let mut language_eces = HashMap::new();
492
493 for (language, group) in language_groups {
494 if group.len() >= 10 {
495 let group_samples: Vec<CalibrationSample> = group.into_iter().cloned().collect();
496 let measurement = self.calculate_calibration_measurement(&group_samples).await?;
497 language_eces.insert(language, measurement.ece);
498 }
499 }
500
501 Ok(language_eces)
502 }
503
504 async fn fit_temperature_model(&self, samples: Vec<&CalibrationSample>, scope: &str) -> Result<TemperatureModel> {
505 let samples_owned: Vec<CalibrationSample> = samples.into_iter().cloned().collect();
507 let before_measurement = self.calculate_calibration_measurement(&samples_owned).await?;
508
509 let mut best_temperature = 1.0;
511 let mut best_ece = before_measurement.ece;
512
513 for temp_candidate in (5..=200).step_by(5) {
514 let temperature = temp_candidate as f32 / 100.0; let calibrated_samples: Vec<CalibrationSample> = samples_owned.iter()
518 .map(|s| {
519 let logit_p = (s.prediction / (1.0 - s.prediction + 1e-8)).ln();
520 let scaled_logit = logit_p / temperature;
521 let calibrated_p = 1.0 / (1.0 + (-scaled_logit).exp());
522
523 CalibrationSample {
524 prediction: calibrated_p.clamp(0.001, 0.999),
525 actual_relevance: s.actual_relevance,
526 query_type: s.query_type.clone(),
527 language: s.language.clone(),
528 }
529 })
530 .collect();
531
532 let temp_measurement = self.calculate_calibration_measurement(&calibrated_samples).await?;
533
534 if temp_measurement.ece < best_ece {
535 best_ece = temp_measurement.ece;
536 best_temperature = temperature;
537 }
538 }
539
540 Ok(TemperatureModel {
541 temperature: best_temperature,
542 bias: 0.0, ece_before: before_measurement.ece,
544 ece_after: best_ece,
545 training_samples: samples_owned.len(),
546 scope: scope.to_string(),
547 })
548 }
549
550 async fn adjust_temperature_models(&self, samples: &[CalibrationSample]) -> Result<()> {
551 info!("Adjusting temperature models due to calibration drift");
552
553 self.train_temperature_models(samples).await?;
555
556 Ok(())
557 }
558
559 async fn add_calibration_snapshot(&self, ece: f32, mce: f32, semantic_active: bool, temperature: f32) {
560 let snapshot = CalibrationSnapshot {
561 timestamp: std::time::SystemTime::now(),
562 ece,
563 mce,
564 sample_count: 0, semantic_features_active: semantic_active,
566 temperature,
567 };
568
569 let mut history = self.calibration_history.write().await;
570 history.push(snapshot);
571
572 let history_len = history.len();
574 if history_len > self.config.measurement_window_size {
575 history.drain(0..history_len - self.config.measurement_window_size);
576 }
577 }
578}
579
580impl Default for CalibrationMeasurement {
583 fn default() -> Self {
584 Self {
585 ece: 0.0,
586 ece_ci_lower: 0.0,
587 ece_ci_upper: 0.0,
588 mce: 0.0,
589 reliability_data: Vec::new(),
590 sample_count: 0,
591 timestamp: std::time::SystemTime::now(),
592 by_query_type: HashMap::new(),
593 by_language: HashMap::new(),
594 }
595 }
596}
597
598impl Default for FeatureScalerSet {
599 fn default() -> Self {
600 Self {
601 semantic_feature_caps: HashMap::new(),
602 linear_scalers: HashMap::new(),
603 robust_scalers: HashMap::new(),
604 }
605 }
606}
607
608#[derive(Debug, Clone)]
610pub enum CalibrationStatus {
611 WithinLimits { ece_drift: f32, current_ece: f32 },
612 DriftDetected { ece_drift: f32, current_ece: f32, baseline_ece: f32 },
613 NoBaseline,
614 InsufficientData,
615}
616
617#[derive(Debug, Clone, Serialize, Deserialize)]
619pub struct CalibrationMetrics {
620 pub current_ece: f32,
621 pub baseline_ece: f32,
622 pub ece_drift: f32,
623 pub max_allowed_drift: f32,
624 pub within_limits: bool,
625 pub recent_trend: f32,
626 pub measurement_count: usize,
627 pub sample_count: usize,
628}
629
630pub async fn initialize_calibration(config: &CalibrationConfig) -> Result<()> {
632 info!("Initializing calibration preservation system");
633 info!("ECE drift limit: {:.4}, log-odds cap: {:.2}",
634 config.max_ece_drift, config.log_odds_cap);
635
636 if config.max_ece_drift > 0.01 {
637 warn!("ECE drift limit {:.4} may be too permissive", config.max_ece_drift);
638 }
639
640 if config.log_odds_cap < 3.0 {
641 warn!("Log-odds cap {:.2} may be too restrictive", config.log_odds_cap);
642 }
643
644 info!("Calibration system initialized");
645 Ok(())
646}
647
648#[cfg(test)]
649mod tests {
650 use super::*;
651
652 #[tokio::test]
653 async fn test_calibration_system_creation() {
654 let config = CalibrationConfig {
655 max_ece_drift: 0.005,
656 log_odds_cap: 5.0,
657 temperature: 1.0,
658 min_samples_for_calibration: 100,
659 measurement_window_size: 50,
660 auto_temperature_adjustment: true,
661 };
662
663 let system = CalibrationSystem::new(config).await.unwrap();
664 let metrics = system.get_calibration_metrics().await;
665 assert_eq!(metrics.measurement_count, 0);
666 }
667
668 #[tokio::test]
669 async fn test_feature_scaling() {
670 let config = CalibrationConfig {
671 max_ece_drift: 0.005,
672 log_odds_cap: 3.0,
673 temperature: 1.0,
674 min_samples_for_calibration: 100,
675 measurement_window_size: 50,
676 auto_temperature_adjustment: false,
677 };
678
679 let system = CalibrationSystem::new(config).await.unwrap();
680
681 let mut features = HashMap::new();
682 features.insert("semantic_similarity".to_string(), 8.0); features.insert("lexical_score".to_string(), 0.7);
684
685 let scaled = system.scale_features(&features).await.unwrap();
686
687 assert!(scaled["semantic_similarity"].abs() <= 3.0);
689 assert_eq!(scaled["lexical_score"], 0.7);
691 }
692
693 #[tokio::test]
694 async fn test_temperature_scaling() {
695 let config = CalibrationConfig {
696 max_ece_drift: 0.005,
697 log_odds_cap: 5.0,
698 temperature: 2.0, min_samples_for_calibration: 100,
700 measurement_window_size: 50,
701 auto_temperature_adjustment: false,
702 };
703
704 let system = CalibrationSystem::new(config).await.unwrap();
705
706 let extreme_prediction = 0.95;
707 let calibrated = system.apply_temperature_scaling(extreme_prediction, "natural_language", None).await.unwrap();
708
709 assert!(calibrated < extreme_prediction);
711 assert!(calibrated > 0.5); }
713
714 #[test]
715 fn test_wilson_confidence_interval() {
716 let config = CalibrationConfig {
717 max_ece_drift: 0.005,
718 log_odds_cap: 5.0,
719 temperature: 1.0,
720 min_samples_for_calibration: 100,
721 measurement_window_size: 50,
722 auto_temperature_adjustment: false,
723 };
724
725 let system = CalibrationSystem {
726 config,
727 baseline_calibration: Arc::new(RwLock::new(None)),
728 current_calibration: Arc::new(RwLock::new(CalibrationMeasurement::default())),
729 temperature_models: Arc::new(RwLock::new(HashMap::new())),
730 feature_scalers: Arc::new(RwLock::new(FeatureScalerSet::default())),
731 calibration_history: Arc::new(RwLock::new(Vec::new())),
732 };
733
734 let (lower, upper) = system.calculate_wilson_confidence_interval(0.5, 100);
735
736 assert!(lower < 0.5);
738 assert!(upper > 0.5);
739 assert!(upper - lower < 0.2); }
741
742 #[tokio::test]
743 async fn test_calibration_measurement() {
744 let config = CalibrationConfig {
745 max_ece_drift: 0.005,
746 log_odds_cap: 5.0,
747 temperature: 1.0,
748 min_samples_for_calibration: 10,
749 measurement_window_size: 50,
750 auto_temperature_adjustment: false,
751 };
752
753 let system = CalibrationSystem::new(config).await.unwrap();
754
755 let samples = vec![
757 CalibrationSample { prediction: 0.1, actual_relevance: 0.0, query_type: "test".to_string(), language: None },
758 CalibrationSample { prediction: 0.3, actual_relevance: 0.0, query_type: "test".to_string(), language: None },
759 CalibrationSample { prediction: 0.5, actual_relevance: 1.0, query_type: "test".to_string(), language: None },
760 CalibrationSample { prediction: 0.7, actual_relevance: 1.0, query_type: "test".to_string(), language: None },
761 CalibrationSample { prediction: 0.9, actual_relevance: 1.0, query_type: "test".to_string(), language: None },
762 ];
763
764 let measurement = system.calculate_calibration_measurement(&samples).await.unwrap();
765
766 assert!(measurement.ece >= 0.0);
767 assert!(measurement.ece <= 1.0);
768 assert!(measurement.mce >= 0.0);
769 assert_eq!(measurement.sample_count, 5);
770 }
771}