use std::collections::{HashMap, HashSet};
use super::criteria::{Dataset, EvalError, Evaluator, PairwiseEvaluator, Predictor, Score};
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ExampleReport {
pub index: usize,
pub input: String,
pub reference: String,
pub prediction: String,
pub scores: HashMap<String, Score>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ScoreSummary {
pub mean: f64,
pub std: f64,
pub count: usize,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct FailureRecord {
pub index: usize,
pub stage: String,
pub error: String,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct Report {
pub per_example: Vec<ExampleReport>,
pub summary: HashMap<String, ScoreSummary>,
pub failures: Vec<FailureRecord>,
}
pub struct EvalRunner {
evaluators: Vec<Box<dyn Evaluator>>,
pairwise: Vec<Box<dyn PairwiseEvaluator>>,
}
impl EvalRunner {
pub fn new(evaluators: Vec<Box<dyn Evaluator>>) -> Self {
Self {
evaluators,
pairwise: Vec::new(),
}
}
pub fn with_pairwise(mut self, pairwise: Vec<Box<dyn PairwiseEvaluator>>) -> Self {
self.pairwise.extend(pairwise);
self
}
pub async fn run(
&self,
dataset: &Dataset,
predictor: &dyn Predictor,
) -> Result<Report, EvalError> {
Self::warn_duplicate_names(&self.evaluators, &self.pairwise);
let mut per_example = Vec::with_capacity(dataset.len());
let mut failures = Vec::new();
let mut per_name: HashMap<String, Vec<f64>> = HashMap::new();
for (i, ex) in dataset.examples.iter().enumerate() {
let prediction = match predictor.predict(&ex.input).await {
Ok(p) => p,
Err(e) => {
failures.push(FailureRecord {
index: i,
stage: "predict".into(),
error: e.to_string(),
});
continue;
}
};
let mut scores = HashMap::new();
for ev in &self.evaluators {
match ev.eval(&ex.input, &prediction, &ex.reference).await {
Ok(s) => {
per_name
.entry(ev.name().to_string())
.or_default()
.push(s.value);
scores.insert(ev.name().to_string(), s);
}
Err(e) => failures.push(FailureRecord {
index: i,
stage: ev.name().to_string(),
error: e.to_string(),
}),
}
}
for ev in &self.pairwise {
match ev.eval_pair(&ex.input, &prediction, &ex.reference).await {
Ok(s) => {
per_name
.entry(ev.name().to_string())
.or_default()
.push(s.value);
scores.insert(ev.name().to_string(), s);
}
Err(e) => failures.push(FailureRecord {
index: i,
stage: ev.name().to_string(),
error: e.to_string(),
}),
}
}
per_example.push(ExampleReport {
index: i,
input: ex.input.clone(),
reference: ex.reference.clone(),
prediction,
scores,
});
}
let mut summary = HashMap::new();
for (name, values) in per_name {
let count = values.len();
let mean = values.iter().sum::<f64>() / count as f64;
let variance = values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / count as f64;
summary.insert(
name,
ScoreSummary {
mean,
std: variance.sqrt(),
count,
},
);
}
Ok(Report {
per_example,
summary,
failures,
})
}
fn warn_duplicate_names(
evaluators: &[Box<dyn Evaluator>],
pairwise: &[Box<dyn PairwiseEvaluator>],
) {
let mut seen = HashSet::new();
for ev in evaluators {
if !seen.insert(ev.name()) {
log::warn!(
"EvalRunner: 存在重名评测器 '{}',报告数据会被覆盖",
ev.name()
);
}
}
for ev in pairwise {
if !seen.insert(ev.name()) {
log::warn!(
"EvalRunner: 存在重名评测器 '{}',报告数据会被覆盖",
ev.name()
);
}
}
}
}