Skip to main content

lc_evaluation/
runner.rs

1//! 批量运行器:`Report` 与 `EvalRunner`。
2//!
3//! `EvalRunner` 在数据集上逐条调用 `Predictor`,再交给多个单点 `Evaluator`
4//! 与成对 `PairwiseEvaluator` 打分,最终汇总为 `Report`。
5//!
6//! P1-3: 逐条容错——单条 predict 或某评测器打分失败记入 `Report::failures`,
7//! 已算好的结果不丢弃,整体不中止。P1-4: `Report` 携带原始文本 + 标准差,
8//! 并实现 `Serialize`/`Deserialize`,便于落盘后二次分析。
9
10use std::collections::{HashMap, HashSet};
11
12use super::criteria::{Dataset, EvalError, Evaluator, PairwiseEvaluator, Predictor, Score};
13
14/// 单条样例的完整评测记录(含原始文本,便于出低分时追溯)。
15#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
16pub struct ExampleReport {
17    /// 样例在数据集中的下标(0 起)
18    pub index: usize,
19    pub input: String,
20    pub reference: String,
21    pub prediction: String,
22    /// 各评测器对该条的得分(失败或未运行的评测器不在其中)
23    pub scores: HashMap<String, Score>,
24}
25
26/// 单个评测器的汇总统计(均值 + 总体标准差 + 样本数)。
27#[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/// 单条失败记录:某下标样例的 predict 或某评测器打分失败。
35#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
36pub struct FailureRecord {
37    /// 样例在数据集中的下标(0 起)
38    pub index: usize,
39    /// 失败阶段:`"predict"` 或评测器 `name()`
40    pub stage: String,
41    pub error: String,
42}
43
44/// 评测报告(含原文、标准差、失败清单;可反序列化)。
45#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
46pub struct Report {
47    /// 逐条完整记录(含 input/reference/prediction 原文)
48    pub per_example: Vec<ExampleReport>,
49    /// 各评测器汇总(均值 + 标准差 + 样本数)
50    pub summary: HashMap<String, ScoreSummary>,
51    /// 逐条容错收集的失败记录(为空表示全部成功)
52    pub failures: Vec<FailureRecord>,
53}
54
55/// 批量运行器:同时收纳单点与成对评测器。
56pub 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    /// 追加成对评测器(P1-1,竞技场评测进统一报告)。
70    pub fn with_pairwise(mut self, pairwise: Vec<Box<dyn PairwiseEvaluator>>) -> Self {
71        self.pairwise.extend(pairwise);
72        self
73    }
74
75    /// 在数据集上运行所有评测器,返回报告。
76    ///
77    /// P1-3: 逐条容错——单条 predict 失败记 `"predict"` 失败记录并跳过该条;
78    /// 某评测器打分失败只记该评测器的失败记录,其它评测器照常出分。
79    /// P1-1: 成对评测器同样参与,以 `(prediction, reference)` 作为 A/B 两个候选
80    /// (竞技场用法:把待比答案放进 reference 槽)。
81    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        // 每个评测器累计所有成功的样本分,用于算均值/标准差
91        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            // 总体标准差:分布/方差信息比均值更能反映评测器的稳定性
154            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    /// P1-4: 重名评测器会在 summary/报告里静默互相覆盖,至少 `log::warn` 提示。
173    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}