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 {
63 Self {
64 evaluators,
65 pairwise: Vec::new(),
66 }
67 }
68
69 pub fn with_pairwise(mut self, pairwise: Vec<Box<dyn PairwiseEvaluator>>) -> Self {
71 self.pairwise.extend(pairwise);
72 self
73 }
74
75 pub async fn run(
82 &self,
83 dataset: &Dataset,
84 predictor: &dyn Predictor,
85 ) -> Result<Report, EvalError> {
86 Self::warn_duplicate_names(&self.evaluators, &self.pairwise);
87
88 let mut per_example = Vec::with_capacity(dataset.len());
89 let mut failures = Vec::new();
90 let mut per_name: HashMap<String, Vec<f64>> = HashMap::new();
92
93 for (i, ex) in dataset.examples.iter().enumerate() {
94 let prediction = match predictor.predict(&ex.input).await {
95 Ok(p) => p,
96 Err(e) => {
97 failures.push(FailureRecord {
98 index: i,
99 stage: "predict".into(),
100 error: e.to_string(),
101 });
102 continue;
103 }
104 };
105
106 let mut scores = HashMap::new();
107 for ev in &self.evaluators {
108 match ev.eval(&ex.input, &prediction, &ex.reference).await {
109 Ok(s) => {
110 per_name
111 .entry(ev.name().to_string())
112 .or_default()
113 .push(s.value);
114 scores.insert(ev.name().to_string(), s);
115 }
116 Err(e) => failures.push(FailureRecord {
117 index: i,
118 stage: ev.name().to_string(),
119 error: e.to_string(),
120 }),
121 }
122 }
123 for ev in &self.pairwise {
124 match ev.eval_pair(&ex.input, &prediction, &ex.reference).await {
125 Ok(s) => {
126 per_name
127 .entry(ev.name().to_string())
128 .or_default()
129 .push(s.value);
130 scores.insert(ev.name().to_string(), s);
131 }
132 Err(e) => failures.push(FailureRecord {
133 index: i,
134 stage: ev.name().to_string(),
135 error: e.to_string(),
136 }),
137 }
138 }
139
140 per_example.push(ExampleReport {
141 index: i,
142 input: ex.input.clone(),
143 reference: ex.reference.clone(),
144 prediction,
145 scores,
146 });
147 }
148
149 let mut summary = HashMap::new();
150 for (name, values) in per_name {
151 let count = values.len();
152 let mean = values.iter().sum::<f64>() / count as f64;
153 let variance = values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / count as f64;
155 summary.insert(
156 name,
157 ScoreSummary {
158 mean,
159 std: variance.sqrt(),
160 count,
161 },
162 );
163 }
164
165 Ok(Report {
166 per_example,
167 summary,
168 failures,
169 })
170 }
171
172 fn warn_duplicate_names(
174 evaluators: &[Box<dyn Evaluator>],
175 pairwise: &[Box<dyn PairwiseEvaluator>],
176 ) {
177 let mut seen = HashSet::new();
178 for ev in evaluators {
179 if !seen.insert(ev.name()) {
180 log::warn!(
181 "EvalRunner: 存在重名评测器 '{}',报告数据会被覆盖",
182 ev.name()
183 );
184 }
185 }
186 for ev in pairwise {
187 if !seen.insert(ev.name()) {
188 log::warn!(
189 "EvalRunner: 存在重名评测器 '{}',报告数据会被覆盖",
190 ev.name()
191 );
192 }
193 }
194 }
195}