use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use crate::Params;
use crate::protocol::RunResult;
use crate::report::is_na;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct TrialAggregate {
pub eval: String,
pub sample: String,
pub target: String,
#[serde(default, skip_serializing_if = "Params::is_empty")]
pub params: Params,
pub key: String,
pub total: usize,
pub scored: usize,
pub passed: usize,
pub failed: usize,
pub na: usize,
pub skipped: usize,
pub pass_rate: f64,
pub mean: f64,
pub std_dev: f64,
}
impl TrialAggregate {
pub fn repeated(&self) -> bool {
self.total > 1
}
pub fn pass_at_k(&self, k: usize) -> f64 {
pass_at_k(self.scored, self.passed, k)
}
}
pub fn pass_at_k(n: usize, c: usize, k: usize) -> f64 {
if n == 0 || k == 0 {
return 0.0;
}
let c = c.min(n); let k = k.min(n);
if n - c < k {
return 1.0;
}
let mut prod = 1.0_f64;
for i in (n - c + 1)..=n {
prod *= 1.0 - (k as f64) / (i as f64);
}
1.0 - prod
}
pub fn aggregate_trials(results: &[RunResult]) -> Vec<TrialAggregate> {
let mut order: Vec<String> = Vec::new();
let mut groups: BTreeMap<String, Vec<&RunResult>> = BTreeMap::new();
for r in results {
let key = r.logical_key();
if !groups.contains_key(&key) {
order.push(key.clone());
}
groups.entry(key).or_default().push(r);
}
order
.into_iter()
.map(|key| {
let group = &groups[&key];
let first = group[0];
let total = group.len();
let mut scored = 0usize;
let mut passed = 0usize;
let mut na = 0usize;
let mut skipped = 0usize;
let mut values: Vec<f64> = Vec::new();
for r in group {
if r.skipped {
skipped += 1;
} else if is_na(r) {
na += 1;
} else {
scored += 1;
if r.passed {
passed += 1;
}
values.push(r.aggregate);
}
}
let mean = if values.is_empty() {
0.0
} else {
values.iter().sum::<f64>() / values.len() as f64
};
let std_dev = if values.is_empty() {
0.0
} else {
let var =
values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / values.len() as f64;
var.sqrt()
};
let pass_rate = if scored == 0 {
0.0
} else {
passed as f64 / scored as f64
};
TrialAggregate {
eval: first.eval.clone(),
sample: first.sample.clone(),
target: first.target.clone(),
params: first.params.clone(),
key,
total,
scored,
passed,
failed: scored - passed,
na,
skipped,
pass_rate,
mean,
std_dev,
}
})
.collect()
}
pub fn has_trials(results: &[RunResult]) -> bool {
aggregate_trials(results).iter().any(|a| a.repeated())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Score;
use crate::protocol::TranscriptSummary;
fn trial(target: &str, trial: usize, trials: usize, passed: bool) -> RunResult {
RunResult {
eval: "e".into(),
sample: "s".into(),
target: target.into(),
params: Default::default(),
trial,
trials,
seed: Some(trial as u64),
input: Vec::new(),
expected: None,
passed,
aggregate: if passed { 1.0 } else { 0.0 },
scores: vec![if passed {
Score::pass("x", "ok")
} else {
Score::fail("x", "no")
}],
transcript: TranscriptSummary::default(),
skipped: false,
}
}
#[test]
fn pass_at_k_matches_reference() {
assert!((pass_at_k(5, 5, 1) - 1.0).abs() < 1e-9);
assert!((pass_at_k(5, 5, 3) - 1.0).abs() < 1e-9);
assert_eq!(pass_at_k(5, 0, 1), 0.0);
assert!((pass_at_k(10, 4, 1) - 0.4).abs() < 1e-9);
assert!((pass_at_k(2, 1, 2) - 1.0).abs() < 1e-9);
assert!((pass_at_k(4, 2, 2) - (1.0 - 1.0 / 6.0)).abs() < 1e-9);
assert_eq!(pass_at_k(0, 0, 1), 0.0);
assert_eq!(pass_at_k(3, 1, 0), 0.0);
}
#[test]
fn groups_trials_by_logical_key() {
let mut results = Vec::new();
for t in 0..4 {
results.push(trial("sim", t, 4, t != 0)); }
for t in 0..4 {
results.push(trial("opus", t, 4, t == 0)); }
let aggs = aggregate_trials(&results);
assert_eq!(aggs.len(), 2);
let sim = aggs.iter().find(|a| a.target == "sim").unwrap();
assert_eq!(sim.total, 4);
assert_eq!(sim.scored, 4);
assert_eq!(sim.passed, 3);
assert_eq!(sim.failed, 1);
assert!((sim.pass_rate - 0.75).abs() < 1e-9);
assert!(sim.repeated());
assert!((sim.pass_at_k(1) - 0.75).abs() < 1e-9);
assert!((sim.pass_at_k(4) - 1.0).abs() < 1e-9);
assert!((sim.mean - 0.75).abs() < 1e-9);
assert!((sim.std_dev - 0.1875_f64.sqrt()).abs() < 1e-9);
let opus = aggs.iter().find(|a| a.target == "opus").unwrap();
assert_eq!(opus.passed, 1);
assert!((opus.pass_rate - 0.25).abs() < 1e-9);
}
#[test]
fn na_and_skipped_trials_excluded_from_denominator() {
let mut na = trial("sim", 1, 3, false);
na.scores = vec![Score::na("judge", "unreachable")];
let mut skip = trial("sim", 2, 3, false);
skip.skipped = true;
let results = vec![trial("sim", 0, 3, true), na, skip];
let aggs = aggregate_trials(&results);
assert_eq!(aggs.len(), 1);
let a = &aggs[0];
assert_eq!(a.total, 3);
assert_eq!(a.scored, 1); assert_eq!(a.passed, 1);
assert_eq!(a.na, 1);
assert_eq!(a.skipped, 1);
assert!((a.pass_rate - 1.0).abs() < 1e-9);
}
#[test]
fn has_trials_only_for_repeated_cases() {
let single = vec![trial("sim", 0, 1, true)];
assert!(!has_trials(&single));
let repeated = vec![trial("sim", 0, 2, true), trial("sim", 1, 2, false)];
assert!(has_trials(&repeated));
}
}