use crate::observation::Observation;
use indexmap::IndexMap;
use ordered_float::OrderedFloat;
use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone, PartialEq)]
pub struct EvaluatorStableWrongRate {
pub evaluator_id: String,
pub n_keep: usize,
pub n_stable_wrong: usize,
pub stable_wrong_rate: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModelStableWrongRate {
pub model_id: String,
pub n_keep: usize,
pub n_stable_wrong: usize,
pub stable_wrong_rate: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct BudgetStableWrongRate {
pub budget: f64,
pub n_keep: usize,
pub n_stable_wrong: usize,
pub stable_wrong_rate: f64,
}
#[derive(Debug, Clone, Default)]
pub struct StableWrongBreakdown {
pub by_evaluator: Vec<EvaluatorStableWrongRate>,
pub by_model: Vec<ModelStableWrongRate>,
pub by_budget: Vec<BudgetStableWrongRate>,
}
pub fn stable_wrong_breakdown(
observations_by_sample: &IndexMap<String, Vec<Observation>>,
keep_sample_ids: &HashSet<String>,
stable_wrong_sample_ids: &HashSet<String>,
) -> StableWrongBreakdown {
let mut evaluator_keep: HashMap<String, usize> = HashMap::new();
let mut evaluator_wrong: HashMap<String, usize> = HashMap::new();
let mut model_keep: HashMap<String, usize> = HashMap::new();
let mut model_wrong: HashMap<String, usize> = HashMap::new();
let mut budget_keep: HashMap<OrderedFloat<f64>, usize> = HashMap::new();
let mut budget_wrong: HashMap<OrderedFloat<f64>, usize> = HashMap::new();
for sample_id in keep_sample_ids {
let Some(obs) = observations_by_sample.get(sample_id) else {
continue;
};
let is_wrong = stable_wrong_sample_ids.contains(sample_id);
let evaluators: HashSet<&str> = obs
.iter()
.filter_map(|o| o.evaluator_id.as_deref())
.collect();
for e in evaluators {
*evaluator_keep.entry(e.to_string()).or_insert(0) += 1;
if is_wrong {
*evaluator_wrong.entry(e.to_string()).or_insert(0) += 1;
}
}
let models: HashSet<&str> = obs.iter().filter_map(|o| o.model_id.as_deref()).collect();
for m in models {
*model_keep.entry(m.to_string()).or_insert(0) += 1;
if is_wrong {
*model_wrong.entry(m.to_string()).or_insert(0) += 1;
}
}
let budgets: HashSet<OrderedFloat<f64>> = obs
.iter()
.filter_map(|o| o.budget.map(OrderedFloat))
.collect();
for b in budgets {
*budget_keep.entry(b).or_insert(0) += 1;
if is_wrong {
*budget_wrong.entry(b).or_insert(0) += 1;
}
}
}
let mut by_evaluator: Vec<EvaluatorStableWrongRate> = evaluator_keep
.into_iter()
.map(|(evaluator_id, n_keep)| {
let n_stable_wrong = evaluator_wrong.get(&evaluator_id).copied().unwrap_or(0);
EvaluatorStableWrongRate {
stable_wrong_rate: n_stable_wrong as f64 / n_keep as f64,
evaluator_id,
n_keep,
n_stable_wrong,
}
})
.collect();
by_evaluator.sort_by(|a, b| {
b.stable_wrong_rate
.partial_cmp(&a.stable_wrong_rate)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.evaluator_id.cmp(&b.evaluator_id))
});
let mut by_model: Vec<ModelStableWrongRate> = model_keep
.into_iter()
.map(|(model_id, n_keep)| {
let n_stable_wrong = model_wrong.get(&model_id).copied().unwrap_or(0);
ModelStableWrongRate {
stable_wrong_rate: n_stable_wrong as f64 / n_keep as f64,
model_id,
n_keep,
n_stable_wrong,
}
})
.collect();
by_model.sort_by(|a, b| {
b.stable_wrong_rate
.partial_cmp(&a.stable_wrong_rate)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.model_id.cmp(&b.model_id))
});
let mut by_budget: Vec<BudgetStableWrongRate> = budget_keep
.into_iter()
.map(|(budget, n_keep)| {
let n_stable_wrong = budget_wrong.get(&budget).copied().unwrap_or(0);
BudgetStableWrongRate {
stable_wrong_rate: n_stable_wrong as f64 / n_keep as f64,
budget: budget.into_inner(),
n_keep,
n_stable_wrong,
}
})
.collect();
by_budget.sort_by(|a, b| {
b.stable_wrong_rate
.partial_cmp(&a.stable_wrong_rate)
.unwrap_or(std::cmp::Ordering::Equal)
.then(
a.budget
.partial_cmp(&b.budget)
.unwrap_or(std::cmp::Ordering::Equal),
)
});
StableWrongBreakdown {
by_evaluator,
by_model,
by_budget,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn obs(sample_id: &str, evaluator_id: Option<&str>, budget: Option<f64>) -> Observation {
Observation {
sample_id: sample_id.into(),
evaluator_id: evaluator_id.map(String::from),
budget,
..Default::default()
}
}
#[test]
fn test_hand_derived_evaluator_and_model_rates() {
let mut observations_by_sample = IndexMap::new();
observations_by_sample.insert(
"s1".to_string(),
vec![Observation {
model_id: Some("mod1".into()),
..obs("s1", Some("m1"), None)
}],
);
observations_by_sample.insert(
"s2".to_string(),
vec![Observation {
model_id: Some("mod1".into()),
..obs("s2", Some("m1"), None)
}],
);
observations_by_sample.insert(
"s3".to_string(),
vec![Observation {
model_id: Some("mod2".into()),
..obs("s3", Some("m2"), None)
}],
);
let keep_sample_ids: HashSet<String> =
["s1", "s2", "s3"].iter().map(|s| s.to_string()).collect();
let stable_wrong_sample_ids: HashSet<String> =
["s1"].iter().map(|s| s.to_string()).collect();
let breakdown = stable_wrong_breakdown(
&observations_by_sample,
&keep_sample_ids,
&stable_wrong_sample_ids,
);
assert_eq!(
breakdown.by_evaluator,
vec![
EvaluatorStableWrongRate {
evaluator_id: "m1".into(),
n_keep: 2,
n_stable_wrong: 1,
stable_wrong_rate: 0.5,
},
EvaluatorStableWrongRate {
evaluator_id: "m2".into(),
n_keep: 1,
n_stable_wrong: 0,
stable_wrong_rate: 0.0,
},
]
);
assert_eq!(
breakdown.by_model,
vec![
ModelStableWrongRate {
model_id: "mod1".into(),
n_keep: 2,
n_stable_wrong: 1,
stable_wrong_rate: 0.5,
},
ModelStableWrongRate {
model_id: "mod2".into(),
n_keep: 1,
n_stable_wrong: 0,
stable_wrong_rate: 0.0,
},
]
);
assert!(breakdown.by_budget.is_empty());
}
#[test]
fn test_hand_derived_budget_rates_with_dedup() {
let mut observations_by_sample = IndexMap::new();
observations_by_sample.insert(
"s1".to_string(),
vec![obs("s1", None, Some(4.0)), obs("s1", None, Some(4.0))],
);
observations_by_sample.insert("s2".to_string(), vec![obs("s2", None, Some(8.0))]);
let keep_sample_ids: HashSet<String> = ["s1", "s2"].iter().map(|s| s.to_string()).collect();
let stable_wrong_sample_ids: HashSet<String> =
["s1"].iter().map(|s| s.to_string()).collect();
let breakdown = stable_wrong_breakdown(
&observations_by_sample,
&keep_sample_ids,
&stable_wrong_sample_ids,
);
assert_eq!(
breakdown.by_budget,
vec![
BudgetStableWrongRate {
budget: 4.0,
n_keep: 1,
n_stable_wrong: 1,
stable_wrong_rate: 1.0,
},
BudgetStableWrongRate {
budget: 8.0,
n_keep: 1,
n_stable_wrong: 0,
stable_wrong_rate: 0.0,
},
]
);
}
#[test]
fn test_non_keep_samples_never_contribute_no_vacuous_rows() {
let mut observations_by_sample = IndexMap::new();
observations_by_sample.insert(
"s_review".to_string(),
vec![obs("s_review", Some("m3"), None)],
);
let keep_sample_ids: HashSet<String> = HashSet::new(); let stable_wrong_sample_ids: HashSet<String> = HashSet::new();
let breakdown = stable_wrong_breakdown(
&observations_by_sample,
&keep_sample_ids,
&stable_wrong_sample_ids,
);
assert!(
breakdown.by_evaluator.is_empty(),
"m3 never appears on a Keep sample, so it should be absent, not a 0/0 row"
);
}
#[test]
fn test_multi_evaluator_sample_contributes_to_both_non_partition() {
let mut observations_by_sample = IndexMap::new();
observations_by_sample.insert(
"s1".to_string(),
vec![obs("s1", Some("m1"), None), obs("s1", Some("m2"), None)],
);
let keep_sample_ids: HashSet<String> = ["s1"].iter().map(|s| s.to_string()).collect();
let stable_wrong_sample_ids: HashSet<String> = HashSet::new();
let breakdown = stable_wrong_breakdown(
&observations_by_sample,
&keep_sample_ids,
&stable_wrong_sample_ids,
);
let total_n_keep: usize = breakdown.by_evaluator.iter().map(|e| e.n_keep).sum();
assert_eq!(
total_n_keep, 2,
"one Keep sample touching 2 evaluators contributes to both rows' n_keep, so the \
sum (2) exceeds the actual number of Keep samples (1) -- not a partition"
);
}
#[test]
fn test_sort_order_ties_break_by_key_ascending() {
let mut observations_by_sample = IndexMap::new();
observations_by_sample.insert("s1".to_string(), vec![obs("s1", Some("zeta"), None)]);
observations_by_sample.insert("s2".to_string(), vec![obs("s2", Some("alpha"), None)]);
let keep_sample_ids: HashSet<String> = ["s1", "s2"].iter().map(|s| s.to_string()).collect();
let stable_wrong_sample_ids: HashSet<String> = HashSet::new();
let breakdown = stable_wrong_breakdown(
&observations_by_sample,
&keep_sample_ids,
&stable_wrong_sample_ids,
);
let ids: Vec<&str> = breakdown
.by_evaluator
.iter()
.map(|e| e.evaluator_id.as_str())
.collect();
assert_eq!(ids, vec!["alpha", "zeta"]);
}
}