1use std::collections::{HashMap, HashSet};
11
12use super::criteria::{Dataset, EvalError, Evaluator, PairwiseEvaluator, Predictor, Score};
13
14#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
16pub struct ExampleReport {
17 pub index: usize,
19 pub input: String,
20 pub reference: String,
21 pub prediction: String,
22 pub scores: HashMap<String, Score>,
24}
25
26#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
28pub struct ScoreSummary {
29 pub mean: f64,
30 pub std: f64,
31 pub count: usize,
32}
33
34#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
36pub struct FailureRecord {
37 pub index: usize,
39 pub stage: String,
41 pub error: String,
42}
43
44#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
46pub struct Report {
47 pub per_example: Vec<ExampleReport>,
49 pub summary: HashMap<String, ScoreSummary>,
51 pub failures: Vec<FailureRecord>,
53}
54
55pub struct EvalRunner {
57 evaluators: Vec<Box<dyn Evaluator>>,
58 pairwise: Vec<Box<dyn PairwiseEvaluator>>,
59}
60
61impl EvalRunner {
62 pub fn new(evaluators: Vec<Box<dyn Evaluator>>) -> Self {
64 Self {
65 evaluators,
66 pairwise: Vec::new(),
67 }
68 }
69
70 pub fn with_pairwise(mut self, pairwise: Vec<Box<dyn PairwiseEvaluator>>) -> Self {
72 self.pairwise.extend(pairwise);
73 self
74 }
75
76 pub async fn run(
83 &self,
84 dataset: &Dataset,
85 predictor: &dyn Predictor,
86 ) -> Result<Report, EvalError> {
87 Self::warn_duplicate_names(&self.evaluators, &self.pairwise);
88
89 let mut per_example = Vec::with_capacity(dataset.len());
90 let mut failures = Vec::new();
91 let mut per_name: HashMap<String, Vec<f64>> = HashMap::new();
93
94 for (i, ex) in dataset.examples.iter().enumerate() {
95 let prediction = match predictor.predict(&ex.input).await {
96 Ok(p) => p,
97 Err(e) => {
98 failures.push(FailureRecord {
99 index: i,
100 stage: "predict".into(),
101 error: e.to_string(),
102 });
103 continue;
104 }
105 };
106
107 let mut scores = HashMap::new();
108 for ev in &self.evaluators {
109 match ev.eval(&ex.input, &prediction, &ex.reference).await {
110 Ok(s) => {
111 per_name
112 .entry(ev.name().to_string())
113 .or_default()
114 .push(s.value);
115 scores.insert(ev.name().to_string(), s);
116 }
117 Err(e) => failures.push(FailureRecord {
118 index: i,
119 stage: ev.name().to_string(),
120 error: e.to_string(),
121 }),
122 }
123 }
124 for ev in &self.pairwise {
125 match ev.eval_pair(&ex.input, &prediction, &ex.reference).await {
126 Ok(s) => {
127 per_name
128 .entry(ev.name().to_string())
129 .or_default()
130 .push(s.value);
131 scores.insert(ev.name().to_string(), s);
132 }
133 Err(e) => failures.push(FailureRecord {
134 index: i,
135 stage: ev.name().to_string(),
136 error: e.to_string(),
137 }),
138 }
139 }
140
141 per_example.push(ExampleReport {
142 index: i,
143 input: ex.input.clone(),
144 reference: ex.reference.clone(),
145 prediction,
146 scores,
147 });
148 }
149
150 let mut summary = HashMap::new();
151 for (name, values) in per_name {
152 let count = values.len();
153 let mean = values.iter().sum::<f64>() / count as f64;
154 let variance = values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / count as f64;
156 summary.insert(
157 name,
158 ScoreSummary {
159 mean,
160 std: variance.sqrt(),
161 count,
162 },
163 );
164 }
165
166 Ok(Report {
167 per_example,
168 summary,
169 failures,
170 })
171 }
172
173 fn warn_duplicate_names(
175 evaluators: &[Box<dyn Evaluator>],
176 pairwise: &[Box<dyn PairwiseEvaluator>],
177 ) {
178 let mut seen = HashSet::new();
179 for ev in evaluators {
180 if !seen.insert(ev.name()) {
181 log::warn!(
182 "EvalRunner: duplicate evaluator name '{}', report data will be overwritten",
183 ev.name()
184 );
185 }
186 }
187 for ev in pairwise {
188 if !seen.insert(ev.name()) {
189 log::warn!(
190 "EvalRunner: duplicate evaluator name '{}', report data will be overwritten",
191 ev.name()
192 );
193 }
194 }
195 }
196}