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>,
#[serde(default)]
pub cost: crate::OverallCost,
}
pub struct EvalRunner {
evaluators: Vec<Box<dyn Evaluator>>,
pairwise: Vec<Box<dyn PairwiseEvaluator>>,
price_book: crate::PriceBook,
}
impl EvalRunner {
pub fn new(evaluators: Vec<Box<dyn Evaluator>>) -> Self {
Self {
evaluators,
pairwise: Vec::new(),
price_book: crate::PriceBook::default_set(),
}
}
pub fn with_pairwise(mut self, pairwise: Vec<Box<dyn PairwiseEvaluator>>) -> Self {
self.pairwise.extend(pairwise);
self
}
pub fn with_price_book(mut self, price_book: crate::PriceBook) -> Self {
self.price_book = price_book;
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();
let mut cost = crate::OverallCost::default();
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;
}
};
if let Some(usage) = predictor.report_token_usage().await {
cost.accumulate(&usage, &self.price_book);
}
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,
cost,
})
}
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: duplicate evaluator name '{}', report data will be overwritten",
ev.name()
);
}
}
for ev in pairwise {
if !seen.insert(ev.name()) {
log::warn!(
"EvalRunner: duplicate evaluator name '{}', report data will be overwritten",
ev.name()
);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
struct ConstantEvaluator;
#[async_trait]
impl Evaluator for ConstantEvaluator {
async fn eval(
&self,
_input: &str,
_prediction: &str,
_reference: &str,
) -> Result<Score, EvalError> {
Ok(Score::new(1.0))
}
fn name(&self) -> &str {
"constant"
}
}
struct UsagePredictor;
#[async_trait]
impl Predictor for UsagePredictor {
async fn predict(&self, input: &str) -> Result<String, EvalError> {
Ok(format!("{input}!"))
}
async fn report_token_usage(&self) -> Option<crate::TokenUsage> {
Some(crate::TokenUsage {
prompt_tokens: 100,
completion_tokens: 50,
model: Some("gpt-4o-mini".into()),
})
}
}
struct UnmeteredPredictor;
#[async_trait]
impl Predictor for UnmeteredPredictor {
async fn predict(&self, input: &str) -> Result<String, EvalError> {
Ok(input.to_string())
}
}
async fn dataset2() -> Dataset {
Dataset::new(vec![
crate::Example::new("q1", "a1"),
crate::Example::new("q2", "a2"),
])
}
#[tokio::test]
async fn old_report_without_cost_field_still_deserializes() {
let old_json = r#"{
"per_example": [],
"summary": {},
"failures": []
}"#;
let report: Report = serde_json::from_str(old_json).unwrap();
assert_eq!(report.cost.total_tokens, 0);
assert!(report.cost.cost_usd.is_none());
}
#[tokio::test]
async fn report_round_trips_with_cost_field() {
let runner = EvalRunner::new(vec![Box::new(ConstantEvaluator)])
.with_price_book(crate::PriceBook::default_set());
let report = runner
.run(&dataset2().await, &UsagePredictor)
.await
.unwrap();
let json = serde_json::to_string(&report).unwrap();
let back: Report = serde_json::from_str(&json).unwrap();
assert_eq!(back.cost.total_tokens, report.cost.total_tokens);
assert_eq!(back.cost.cost_usd, report.cost.cost_usd);
}
#[tokio::test]
async fn runner_accumulates_priced_usage_across_examples() {
let runner = EvalRunner::new(vec![Box::new(ConstantEvaluator)])
.with_price_book(crate::PriceBook::default_set());
let report = runner
.run(&dataset2().await, &UsagePredictor)
.await
.unwrap();
assert_eq!(report.cost.prompt_tokens, 200);
assert_eq!(report.cost.completion_tokens, 100);
assert_eq!(report.cost.total_tokens, 300);
let expected = 200.0 / 1e6 * 0.15 + 100.0 / 1e6 * 0.60; let usd = report.cost.cost_usd.unwrap();
assert!(
(usd - expected).abs() < 1e-12,
"got {usd}, expected {expected}"
);
}
#[tokio::test]
async fn unmetered_predictor_leaves_cost_at_zero() {
let runner = EvalRunner::new(vec![Box::new(ConstantEvaluator)]);
let report = runner
.run(&dataset2().await, &UnmeteredPredictor)
.await
.unwrap();
assert_eq!(report.cost.total_tokens, 0);
assert!(report.cost.cost_usd.is_none());
}
}