use crate::config::{DecisionScore, ScoreConfig};
use crate::observation::Observation;
use crate::schema::StabilityReport;
use crate::score_all_gold_weighted;
use crate::scoring::score_all;
#[derive(Debug)]
pub struct CalibrationResult {
pub keep_threshold: f64,
pub drop_threshold: f64,
pub decision_score_name: &'static str,
pub achieved_precision: f64,
pub coverage: f64,
pub n_keep: usize,
pub n_total: usize,
}
pub fn compute_calibration(
observations: &[Observation],
decision_score: &DecisionScore,
confidence_level: f64,
target_precision: f64,
target_coverage: Option<f64>,
drop_threshold: f64,
) -> Option<CalibrationResult> {
let gold: std::collections::HashMap<&str, &str> = observations
.iter()
.filter_map(|o| o.gold_label.as_deref().map(|g| (o.sample_id.as_str(), g)))
.collect();
if gold.is_empty() {
return None;
}
let config = ScoreConfig {
thresholds: crate::decision::Thresholds {
keep: 0.0,
drop: -1.0,
},
decision_score: decision_score.clone(),
confidence_level,
..ScoreConfig::default()
};
let reports = if matches!(decision_score, DecisionScore::GoldWeighted) {
score_all_gold_weighted(observations.to_vec(), &config)
} else {
score_all(observations.to_vec(), &config)
};
let n_total = reports.len();
if n_total == 0 {
return None;
}
let decision_score_name = match decision_score {
DecisionScore::Raw => "raw",
DecisionScore::Adjusted => "adjusted",
DecisionScore::LowerConfidenceBound => "lcb",
DecisionScore::GoldWeighted => "gold-weighted",
};
for i in 0..=49usize {
let t = 0.99 - i as f64 * 0.01;
let score_val = |r: &StabilityReport| match decision_score {
DecisionScore::Adjusted => r.adjusted_stability_score,
_ => r.stability_score,
};
let kept: Vec<&StabilityReport> = reports.iter().filter(|r| score_val(r) >= t).collect();
if kept.is_empty() {
continue;
}
let n_keep = kept.len();
let coverage = n_keep as f64 / n_total as f64;
if let Some(tc) = target_coverage
&& coverage < tc
{
continue;
}
let matches = kept
.iter()
.filter(|r| {
let label = r
.weighted_majority_label
.as_deref()
.or(r.majority_label.as_deref());
gold.get(r.sample_id.as_str())
.and_then(|&g| label.map(|m| m == g))
.unwrap_or(false)
})
.count();
let precision = matches as f64 / n_keep as f64;
if precision >= target_precision {
return Some(CalibrationResult {
keep_threshold: t,
drop_threshold,
decision_score_name,
achieved_precision: precision,
coverage,
n_keep,
n_total,
});
}
}
None
}