1use super::task::{EvalRunResult, RunOutcome};
2use serde::{Deserialize, Serialize};
3
4#[derive(Debug, Clone, Serialize, Deserialize)]
5pub struct EvalMetric {
6 pub(crate) pass_at_k: f64,
7 pub(crate) pass_power_k: f64,
8 pub(crate) pass_all_k: f64,
9 pub(crate) k: u32,
10 pub(crate) total_runs: u32,
11 pub(crate) passed_runs: u32,
12 pub(crate) task_id: String,
13}
14
15pub fn compute_metric(task_id: &str, results: &[EvalRunResult]) -> EvalMetric {
16 compute_metric_with_k(task_id, results, 1).unwrap_or_else(|_| empty_metric(task_id))
17}
18
19pub fn compute_metric_with_k(task_id: &str, results: &[EvalRunResult], k: u32) -> anyhow::Result<EvalMetric> {
21 let total = results.len() as u32;
22 anyhow::ensure!(total > 0, "cannot compute evaluation metrics without attempts");
23 anyhow::ensure!(k > 0 && k <= total, "metric k must be between 1 and the number of attempts");
24 let passed = results.iter().filter(|r| r.outcome == RunOutcome::Pass).count() as u32;
25 Ok(EvalMetric {
26 pass_at_k: pass_at_k_counts(total, passed, k),
27 pass_power_k: (passed as f64 / total as f64).powi(k as i32),
28 pass_all_k: if passed == total { 1.0 } else { 0.0 },
29 k,
30 total_runs: total,
31 passed_runs: passed,
32 task_id: task_id.into(),
33 })
34}
35
36fn empty_metric(task_id: &str) -> EvalMetric {
37 EvalMetric {
38 pass_at_k: 0.0,
39 pass_power_k: 0.0,
40 pass_all_k: 0.0,
41 k: 0,
42 total_runs: 0,
43 passed_runs: 0,
44 task_id: task_id.into(),
45 }
46}
47
48pub fn aggregate_metrics(metrics: &[EvalMetric]) -> EvalMetric {
49 if metrics.is_empty() {
50 return EvalMetric {
51 pass_at_k: 0.0,
52 pass_power_k: 0.0,
53 pass_all_k: 0.0,
54 k: 0,
55 total_runs: 0,
56 passed_runs: 0,
57 task_id: "aggregate".into(),
58 };
59 }
60 let total_runs: u32 = metrics.iter().map(|m| m.total_runs).sum();
61 let passed_runs: u32 = metrics.iter().map(|m| m.passed_runs).sum();
62 let pass_at_k = metrics.iter().map(|metric| metric.pass_at_k).sum::<f64>() / metrics.len() as f64;
63 let pass_power_k = metrics.iter().map(|metric| metric.pass_power_k).sum::<f64>() / metrics.len() as f64;
64 let pass_all_k = if metrics.iter().all(|m| m.pass_all_k > 0.0) {
65 1.0
66 } else {
67 0.0
68 };
69 EvalMetric {
70 pass_at_k,
71 pass_power_k,
72 pass_all_k,
73 k: metrics.iter().map(|metric| metric.k).max().unwrap_or(0),
74 total_runs,
75 passed_runs,
76 task_id: "aggregate".into(),
77 }
78}
79
80pub fn pass_at_k(results: &[EvalRunResult]) -> f64 {
81 pass_at_k_with_k(results, 1).unwrap_or(0.0)
82}
83
84pub fn pass_at_k_with_k(results: &[EvalRunResult], k: u32) -> anyhow::Result<f64> {
86 let total = u32::try_from(results.len()).map_err(|_| anyhow::anyhow!("too many evaluation attempts"))?;
87 let passed = u32::try_from(results.iter().filter(|result| result.outcome == RunOutcome::Pass).count())
88 .map_err(|_| anyhow::anyhow!("too many passed attempts"))?;
89 anyhow::ensure!(total > 0 && k > 0 && k <= total, "pass@k requires 1 <= k <= attempts");
90 Ok(pass_at_k_counts(total, passed, k))
91}
92
93pub fn pass_power_k(results: &[EvalRunResult], k: u32) -> anyhow::Result<f64> {
95 let total = results.len() as u32;
96 anyhow::ensure!(total > 0 && k > 0, "pass^k requires at least one attempt and k");
97 let passed = results.iter().filter(|result| result.outcome == RunOutcome::Pass).count() as f64;
98 Ok((passed / total as f64).powi(k as i32))
99}
100
101fn pass_at_k_counts(total: u32, passed: u32, k: u32) -> f64 {
102 1.0 - combinations(total.saturating_sub(passed), k) / combinations(total, k)
103}
104
105fn combinations(n: u32, k: u32) -> f64 {
106 if k > n {
107 return 0.0;
108 }
109 let k = k.min(n - k);
110 let mut result = 1.0;
111 for i in 1..=k {
112 result *= f64::from(n - k + i) / f64::from(i);
113 }
114 result
115}
116
117pub fn pass_all_k(results: &[EvalRunResult]) -> f64 {
118 if results.is_empty() {
119 return 0.0;
120 }
121 if results.iter().all(|r| r.outcome == RunOutcome::Pass) {
122 1.0
123 } else {
124 0.0
125 }
126}
127
128#[cfg(test)]
129mod tests {
130 use super::*;
131 use crate::task::{EvalRunResult, RunOutcome};
132
133 fn r(outcome: RunOutcome) -> EvalRunResult {
134 EvalRunResult {
135 task_id: "t".into(),
136 outcome,
137 error_message: None,
138 duration_secs: 0.0,
139 attempt: 1,
140 cost_usd: None,
141 transcript_path: None,
142 trace_summary: None,
143 }
144 }
145
146 #[test]
147 fn compute_metric_pass_at_k() {
148 let results = vec![r(RunOutcome::Pass), r(RunOutcome::Fail), r(RunOutcome::Error)];
149 let m = compute_metric("t", &results);
150 assert_eq!(m.total_runs, 3);
151 assert_eq!(m.passed_runs, 1);
152 assert!((m.pass_at_k - 1.0 / 3.0).abs() < 1e-9);
153 assert_eq!(m.pass_all_k, 0.0);
154 }
155
156 #[test]
157 fn compute_metric_all_pass() {
158 let results = vec![r(RunOutcome::Pass), r(RunOutcome::Pass)];
159 let m = compute_metric("t", &results);
160 assert_eq!(m.pass_all_k, 1.0);
161 assert!((m.pass_at_k - 1.0).abs() < 1e-9);
162 }
163
164 #[test]
165 fn compute_metric_empty() {
166 let m = compute_metric("t", &[]);
167 assert_eq!(m.total_runs, 0);
168 assert_eq!(m.pass_at_k, 0.0);
169 }
170
171 #[test]
172 fn aggregate_combines_runs() {
173 let a = EvalMetric {
174 pass_at_k: 1.0,
175 pass_power_k: 1.0,
176 pass_all_k: 1.0,
177 k: 1,
178 total_runs: 1,
179 passed_runs: 1,
180 task_id: "a".into(),
181 };
182 let b = EvalMetric {
183 pass_at_k: 0.0,
184 pass_power_k: 0.0,
185 pass_all_k: 0.0,
186 k: 1,
187 total_runs: 1,
188 passed_runs: 0,
189 task_id: "b".into(),
190 };
191 let agg = aggregate_metrics(&[a, b]);
192 assert_eq!(agg.total_runs, 2);
193 assert_eq!(agg.passed_runs, 1);
194 assert!((agg.pass_at_k - 0.5).abs() < 1e-9);
195 assert_eq!(agg.pass_all_k, 0.0);
196 }
197
198 #[test]
199 fn aggregate_empty() {
200 let agg = aggregate_metrics(&[]);
201 assert_eq!(agg.total_runs, 0);
202 assert_eq!(agg.pass_at_k, 0.0);
203 }
204
205 #[test]
206 fn combinatorial_pass_at_k_and_independent_pass_power_are_distinct() {
207 let results = vec![
208 r(RunOutcome::Pass),
209 r(RunOutcome::Fail),
210 r(RunOutcome::Pass),
211 r(RunOutcome::Fail),
212 ];
213 let metric = compute_metric_with_k("t", &results, 2).expect("metric");
214 assert!((metric.pass_at_k - (5.0 / 6.0)).abs() < 1e-9);
215 assert!((metric.pass_power_k - 0.25).abs() < 1e-9);
216 assert!((pass_at_k_with_k(&results, 2).expect("pass@k") - 5.0 / 6.0).abs() < 1e-9);
217 }
218
219 #[test]
220 fn metric_rejects_zero_attempts_and_invalid_k() {
221 assert!(compute_metric_with_k("t", &[], 1).is_err());
222 let results = vec![r(RunOutcome::Pass)];
223 assert!(compute_metric_with_k("t", &results, 0).is_err());
224 assert!(pass_at_k_with_k(&results, 2).is_err());
225 }
226}