Skip to main content

lc_evaluation/
runner.rs

1//! 批量运行器:`Report` 与 `EvalRunner`。
2//!
3//! `EvalRunner` 在数据集上逐条调用 `Predictor`,再交给多个 `Evaluator` 打分,
4//! 最终汇总为 `Report`(逐条得分 + 各评测器均值)。
5
6use std::collections::HashMap;
7
8use super::criteria::{Dataset, EvalError, Evaluator, Predictor, Score};
9
10/// 评测报告
11#[derive(Debug, Clone, serde::Serialize)]
12pub struct Report {
13    /// 每条样例的各评测器得分
14    pub per_example: Vec<HashMap<String, Score>>,
15    /// 各评测器的平均分
16    pub summary: HashMap<String, f64>,
17}
18
19/// 批量运行器
20pub struct EvalRunner {
21    evaluators: Vec<Box<dyn Evaluator>>,
22}
23
24impl EvalRunner {
25    pub fn new(evaluators: Vec<Box<dyn Evaluator>>) -> Self {
26        Self { evaluators }
27    }
28
29    /// 在数据集上运行所有评测器,返回报告
30    pub async fn run(
31        &self,
32        dataset: &Dataset,
33        predictor: &dyn Predictor,
34    ) -> Result<Report, EvalError> {
35        let mut per_example = Vec::with_capacity(dataset.len());
36        let mut sums: HashMap<String, (f64, usize)> = HashMap::new();
37
38        for ex in &dataset.examples {
39            let prediction = predictor.predict(&ex.input).await?;
40            let mut row = HashMap::new();
41            for ev in &self.evaluators {
42                let score = ev.eval(&ex.input, &prediction, &ex.reference).await?;
43                let entry = sums.entry(ev.name().to_string()).or_insert((0.0, 0));
44                entry.0 += score.value;
45                entry.1 += 1;
46                row.insert(ev.name().to_string(), score);
47            }
48            per_example.push(row);
49        }
50
51        let mut summary = HashMap::new();
52        for (name, (total, count)) in sums {
53            if count > 0 {
54                summary.insert(name, total / count as f64);
55            }
56        }
57        Ok(Report {
58            per_example,
59            summary,
60        })
61    }
62}