use crate::build::{PriorAccumulator, SequenceContextTracker};
use crate::error::{Result, Warning};
use crate::input::parse_line;
use crate::model::{BuildConfig, Observation, Outcome, PriorBook, outcome_credit};
use crate::score::ratio;
use serde::Serialize;
use std::collections::{HashMap, HashSet};
use std::io::{BufRead, BufReader, Read};
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
use crate::model::PriorAction;
fn obs(state: &str, action: &str) -> Observation {
Observation {
sequence_id: "seq".to_string(),
step: 0,
state: state.to_string(),
action: action.to_string(),
outcome: Outcome::Success,
score: None,
weight: 1.0,
tags: Vec::new(),
observed_at_unix_seconds: None,
source: None,
}
}
fn sample_book() -> PriorBook {
let mut entries = HashMap::new();
entries.insert(
"s".to_string(),
vec![
PriorAction {
action: "a".into(),
count: 10,
weighted_count: 10.0,
success_rate: Some(0.9),
mean_score: None,
prior: 0.6,
confidence: 0.5,
},
PriorAction {
action: "b".into(),
count: 5,
weighted_count: 5.0,
success_rate: Some(0.5),
mean_score: None,
prior: 0.3,
confidence: 0.3,
},
PriorAction {
action: "c".into(),
count: 2,
weighted_count: 2.0,
success_rate: Some(0.2),
mean_score: None,
prior: 0.1,
confidence: 0.1,
},
],
);
PriorBook {
entries,
..Default::default()
}
}
#[test]
fn is_train_is_stable_for_a_fixed_id() {
assert_eq!(is_train("seq-1", 0.8), is_train("seq-1", 0.8));
}
#[test]
fn is_train_ratio_zero_is_always_false() {
for i in 0..50 {
assert!(!is_train(&format!("seq-{i}"), 0.0));
}
}
#[test]
fn is_train_ratio_one_is_always_true() {
for i in 0..50 {
assert!(is_train(&format!("seq-{i}"), 1.0));
}
}
#[test]
fn is_train_splits_roughly_by_ratio() {
let train_count = (0..1000)
.filter(|i| is_train(&format!("seq-{i}"), 0.8))
.count();
assert!(
(700..=900).contains(&train_count),
"train_count = {train_count}, expected roughly 800/1000"
);
}
#[test]
fn topk_hit_rate_and_mrr_match_hand_computed_ranks() {
let book = sample_book();
let top_k = vec![1, 2, 3];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &[]);
acc.observe(&book, &obs("s", "a")).unwrap(); acc.observe(&book, &obs("s", "b")).unwrap(); acc.observe(&book, &obs("s", "c")).unwrap(); acc.observe(&book, &obs("s", "z")).unwrap();
let report = acc.finish(0);
assert_eq!(report.num_evaluated_observations, 4);
assert_eq!(report.top1_hit_rate, Some(0.25));
assert_eq!(
report.topk_hit_rate,
vec![
TopKHitRate {
k: 1,
hit_rate: Some(0.25)
},
TopKHitRate {
k: 2,
hit_rate: Some(0.5)
},
TopKHitRate {
k: 3,
hit_rate: Some(0.75)
},
]
);
let expected_mrr = (1.0 + 0.5 + 1.0 / 3.0) / 4.0;
assert!((report.mean_reciprocal_rank.unwrap() - expected_mrr).abs() < 1e-9);
assert_eq!(report.avg_rank_when_found, Some(2.0)); }
#[test]
fn unseen_state_counts_as_fallback_not_evaluated() {
let book = sample_book(); let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &[]);
acc.observe(&book, &obs("unseen_state", "a")).unwrap();
let report = acc.finish(0);
assert_eq!(report.num_test_states, 1);
assert_eq!(report.num_fallback_observations, 1);
assert_eq!(report.num_evaluated_observations, 0);
assert_eq!(report.num_test_states_with_candidates, 0);
assert_eq!(report.coverage, Some(0.0));
assert_eq!(report.fallback_rate, Some(1.0));
assert_eq!(report.top1_hit_rate, None);
}
#[test]
fn success_weighted_metrics_equal_unweighted_when_everything_succeeds() {
let book = sample_book();
let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &[]);
acc.observe(&book, &obs("s", "a")).unwrap(); acc.observe(&book, &obs("s", "b")).unwrap(); acc.observe(&book, &obs("s", "z")).unwrap();
let report = acc.finish(0);
assert_eq!(report.success_weighted_top1_hit_rate, report.top1_hit_rate);
assert_eq!(
report.success_weighted_mean_reciprocal_rank,
report.mean_reciprocal_rank
);
}
#[test]
fn success_weighted_metrics_are_none_when_nothing_earns_credit() {
let book = sample_book();
let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &[]);
acc.observe(
&book,
&Observation {
outcome: Outcome::Failure,
..obs("s", "a")
},
)
.unwrap();
acc.observe(
&book,
&Observation {
outcome: Outcome::Unknown,
..obs("s", "a")
},
)
.unwrap();
let report = acc.finish(0);
assert_eq!(report.num_evaluated_observations, 2);
assert_eq!(report.success_weighted_top1_hit_rate, None);
assert_eq!(report.success_weighted_mean_reciprocal_rank, None);
}
#[test]
fn success_weighted_metrics_give_draws_partial_credit() {
let book = sample_book();
let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &[]);
acc.observe(&book, &obs("s", "a")).unwrap(); acc.observe(
&book,
&Observation {
outcome: Outcome::Draw,
..obs("s", "b")
},
)
.unwrap(); acc.observe(
&book,
&Observation {
outcome: Outcome::Failure,
..obs("s", "z")
},
)
.unwrap();
let report = acc.finish(0);
let expected_top1 = 1.0 / 1.5;
let expected_mrr = 1.25 / 1.5;
assert!((report.success_weighted_top1_hit_rate.unwrap() - expected_top1).abs() < 1e-9);
assert!(
(report.success_weighted_mean_reciprocal_rank.unwrap() - expected_mrr).abs() < 1e-9
);
}
#[test]
fn draw_value_zero_makes_success_weighted_metrics_treat_draws_as_failures() {
let book = sample_book();
let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.0, 0, 0, &[]);
acc.observe(
&book,
&Observation {
outcome: Outcome::Draw,
..obs("s", "a")
},
)
.unwrap();
let report = acc.finish(0);
assert_eq!(report.success_weighted_top1_hit_rate, None);
assert_eq!(report.success_weighted_mean_reciprocal_rank, None);
}
#[test]
fn failure_agreement_top1_hit_rate_flags_the_prior_recommending_a_loser() {
let book = sample_book();
let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &[]);
acc.observe(
&book,
&Observation {
outcome: Outcome::Failure,
..obs("s", "a") },
)
.unwrap();
acc.observe(
&book,
&Observation {
outcome: Outcome::Failure,
..obs("s", "b") },
)
.unwrap();
let report = acc.finish(0);
assert_eq!(report.failure_agreement_top1_hit_rate, Some(0.5));
}
#[test]
fn failure_agreement_top1_hit_rate_is_none_without_failure_observations() {
let book = sample_book();
let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &[]);
acc.observe(&book, &obs("s", "a")).unwrap();
let report = acc.finish(0);
assert_eq!(report.failure_agreement_top1_hit_rate, None);
}
#[test]
fn evaluate_end_to_end_matches_hand_derived_expectations() {
let train_ratio = 0.5;
let candidate_ids: Vec<String> = (0..40).map(|i| format!("seq-{i}")).collect();
let train_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| is_train(id, train_ratio))
.collect();
let test_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| !is_train(id, train_ratio))
.collect();
assert!(!train_ids.is_empty(), "need at least one train sequence");
assert!(test_ids.len() >= 2, "need at least two test sequences");
let mut jsonl = String::new();
for id in &train_ids {
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":0,\"state\":\"s\",\"action\":\"a\",\"outcome\":\"success\"}}\n"
));
}
for (i, id) in test_ids.iter().enumerate() {
let action = if i % 2 == 0 { "a" } else { "z" };
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":0,\"state\":\"s\",\"action\":\"{action}\",\"outcome\":\"success\"}}\n"
));
}
let hits = test_ids
.iter()
.enumerate()
.filter(|(i, _)| i % 2 == 0)
.count();
let expected_confidence = train_ids.len() as f64 / (train_ids.len() as f64 + 20.0);
let eval_config = EvalConfig {
train_ratio,
top_k: vec![1],
..EvalConfig::default()
};
let output = evaluate(
jsonl.as_bytes(),
jsonl.as_bytes(),
true,
&BuildConfig::default(),
&eval_config,
)
.unwrap();
assert_eq!(output.report.num_train_observations, train_ids.len() as u64);
assert_eq!(output.report.num_test_observations, test_ids.len() as u64);
assert_eq!(output.report.num_test_states, 1);
assert_eq!(
output.report.num_evaluated_observations,
test_ids.len() as u64
);
assert_eq!(output.report.num_fallback_observations, 0);
assert_eq!(output.report.coverage, Some(1.0));
assert_eq!(output.report.fallback_rate, Some(0.0));
let expected_rate = Some(hits as f64 / test_ids.len() as f64);
assert_eq!(output.report.top1_hit_rate, expected_rate);
assert_eq!(
output.report.topk_hit_rate,
vec![TopKHitRate {
k: 1,
hit_rate: expected_rate
}]
);
assert_eq!(output.report.mean_reciprocal_rank, expected_rate);
assert_eq!(output.report.avg_rank_when_found, Some(1.0));
assert_eq!(
output.report.avg_confidence_on_hit,
Some(expected_confidence)
);
assert_eq!(
output.report.avg_confidence_on_miss,
Some(expected_confidence)
);
assert_eq!(output.report.score_lift, None); assert!(output.warnings.is_empty());
}
#[test]
fn evaluate_reports_context_aware_metrics_when_context_order_is_set() {
let train_ratio = 0.5;
let candidate_ids: Vec<String> = (0..60).map(|i| format!("seq-{i}")).collect();
let train_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| is_train(id, train_ratio))
.collect();
let test_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| !is_train(id, train_ratio))
.collect();
assert!(train_ids.len() >= 2, "need at least two train sequences");
assert!(test_ids.len() >= 2, "need at least two test sequences");
let mut jsonl = String::new();
for (i, id) in train_ids.iter().chain(test_ids.iter()).enumerate() {
let (first, second) = if i % 2 == 0 { ("x", "A") } else { ("y", "B") };
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":0,\"state\":\"s0\",\"action\":\"{first}\",\"outcome\":\"success\"}}\n"
));
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":1,\"state\":\"s\",\"action\":\"{second}\",\"outcome\":\"success\"}}\n"
));
}
let build_config = BuildConfig {
context_order: 1,
..Default::default()
};
let eval_config = EvalConfig {
train_ratio,
top_k: vec![1],
..EvalConfig::default()
};
let output = evaluate(
jsonl.as_bytes(),
jsonl.as_bytes(),
true,
&build_config,
&eval_config,
)
.unwrap();
let report = output.report;
assert!(report.context_top1_hit_rate.unwrap() > report.top1_hit_rate.unwrap());
assert!(
report.context_mean_reciprocal_rank.unwrap() > report.mean_reciprocal_rank.unwrap()
);
let order1 = report
.hit_rate_by_matched_order
.iter()
.find(|m| m.order == 1)
.expect("order-1 entries should exist");
assert_eq!(order1.top1_hit_rate, Some(1.0));
assert_eq!(order1.num_evaluated, test_ids.len() as u64);
}
#[test]
fn evaluate_context_fields_are_absent_when_context_order_is_zero() {
let train_ratio = 0.5;
let candidate_ids: Vec<String> = (0..40).map(|i| format!("seq-{i}")).collect();
let train_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| is_train(id, train_ratio))
.collect();
let test_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| !is_train(id, train_ratio))
.collect();
assert!(!train_ids.is_empty());
assert!(!test_ids.is_empty());
let mut jsonl = String::new();
for id in train_ids.iter().chain(test_ids.iter()) {
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":0,\"state\":\"s\",\"action\":\"a\",\"outcome\":\"success\"}}\n"
));
}
let eval_config = EvalConfig {
train_ratio,
top_k: vec![1],
..EvalConfig::default()
};
let output = evaluate(
jsonl.as_bytes(),
jsonl.as_bytes(),
true,
&BuildConfig::default(),
&eval_config,
)
.unwrap();
assert!(output.report.num_evaluated_observations > 0);
assert_eq!(output.report.context_top1_hit_rate, None);
assert_eq!(output.report.context_mean_reciprocal_rank, None);
assert!(output.report.hit_rate_by_matched_order.is_empty());
}
#[test]
fn evaluate_applies_time_decay_to_its_train_side_prior_same_as_build() {
let train_ratio = 0.5;
let candidate_ids: Vec<String> = (0..40).map(|i| format!("seq-{i}")).collect();
let train_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| is_train(id, train_ratio))
.collect();
let test_ids: Vec<&String> = candidate_ids
.iter()
.filter(|id| !is_train(id, train_ratio))
.collect();
assert!(train_ids.len() >= 2, "need at least two train sequences");
assert!(!test_ids.is_empty(), "need at least one test sequence");
let split = (train_ids.len() * 4 / 5).clamp(1, train_ids.len() - 1);
let reference: i64 = 1_000_000;
let mut jsonl = String::new();
for id in &train_ids[..split] {
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":0,\"state\":\"s\",\"action\":\"old_winner\",\
\"outcome\":\"success\",\"observed_at_unix_seconds\":{}}}\n",
reference - 5000 * 86_400, ));
}
for id in &train_ids[split..] {
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":0,\"state\":\"s\",\"action\":\"new_winner\",\
\"outcome\":\"success\",\"observed_at_unix_seconds\":{reference}}}\n"
));
}
for id in &test_ids {
jsonl.push_str(&format!(
"{{\"sequence_id\":\"{id}\",\"step\":0,\"state\":\"s\",\"action\":\"new_winner\",\"outcome\":\"success\"}}\n"
));
}
let eval_config = EvalConfig {
train_ratio,
top_k: vec![1],
..EvalConfig::default()
};
let without_decay = evaluate(
jsonl.as_bytes(),
jsonl.as_bytes(),
true,
&BuildConfig::default(),
&eval_config,
)
.unwrap();
assert_eq!(
without_decay.report.top1_hit_rate,
Some(0.0),
"without decay, old_winner's raw count should dominate and mismatch every \
new_winner test observation"
);
let decay_config = BuildConfig {
time_decay_half_life_days: Some(5.0),
time_decay_reference_unix_seconds: Some(reference),
..Default::default()
};
let with_decay = evaluate(
jsonl.as_bytes(),
jsonl.as_bytes(),
true,
&decay_config,
&eval_config,
)
.unwrap();
assert_eq!(
with_decay.report.top1_hit_rate,
Some(1.0),
"with decay, old_winner's effective weight should be crushed, flipping the \
#1 ranking to new_winner"
);
}
fn calibration_fixture_book() -> PriorBook {
let action_with_confidence = |confidence: f64| PriorAction {
action: "a".into(),
count: 1,
weighted_count: 1.0,
success_rate: None,
mean_score: None,
prior: 1.0,
confidence,
};
let mut entries = HashMap::new();
entries.insert("s1".to_string(), vec![action_with_confidence(0.05)]);
entries.insert("s2".to_string(), vec![action_with_confidence(0.55)]);
entries.insert("s3".to_string(), vec![action_with_confidence(1.0)]);
PriorBook {
entries,
..Default::default()
}
}
#[test]
fn calibration_bins_are_deterministic_length_and_bucketed_correctly() {
let book = calibration_fixture_book();
let top_k = vec![1];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 10, &[]);
acc.observe(&book, &obs("s1", "a")).unwrap(); acc.observe(&book, &obs("s2", "b")).unwrap(); acc.observe(&book, &obs("s3", "a")).unwrap();
let report = acc.finish(0);
assert_eq!(report.confidence_calibration.len(), 10);
let bin0 = &report.confidence_calibration[0];
assert!((bin0.min_confidence - 0.0).abs() < 1e-9);
assert!((bin0.max_confidence - 0.1).abs() < 1e-9);
assert_eq!(bin0.num_evaluated, 1);
assert_eq!(bin0.top1_hit_rate, Some(1.0));
assert_eq!(bin0.mean_reciprocal_rank, Some(1.0));
let bin5 = &report.confidence_calibration[5];
assert_eq!(bin5.num_evaluated, 1);
assert_eq!(bin5.top1_hit_rate, Some(0.0));
assert_eq!(bin5.mean_reciprocal_rank, Some(0.0));
let bin9 = &report.confidence_calibration[9]; assert_eq!(bin9.num_evaluated, 1);
assert_eq!(bin9.top1_hit_rate, Some(1.0));
let bin1 = &report.confidence_calibration[1];
assert_eq!(bin1.num_evaluated, 0);
assert_eq!(bin1.top1_hit_rate, None);
assert_eq!(bin1.mean_reciprocal_rank, None);
}
#[test]
fn threshold_sweep_matches_hand_computed_fixture() {
let book = calibration_fixture_book();
let top_k = vec![1];
let thresholds = vec![0.1, 0.6];
let mut acc = EvalAccumulator::new(&top_k, 0.5, 0, 0, &thresholds);
acc.observe(&book, &obs("s1", "a")).unwrap(); acc.observe(&book, &obs("s2", "b")).unwrap(); acc.observe(&book, &obs("s3", "a")).unwrap();
let report = acc.finish(0);
assert_eq!(report.threshold_sweep.len(), 2);
let at_0_1 = &report.threshold_sweep[0];
assert_eq!(at_0_1.min_confidence, 0.1);
assert!((at_0_1.covered_fraction - 2.0 / 3.0).abs() < 1e-9); assert!((at_0_1.abstained_fraction - 1.0 / 3.0).abs() < 1e-9);
assert_eq!(at_0_1.top1_hit_rate, Some(0.5)); assert_eq!(at_0_1.mean_reciprocal_rank, Some(0.5));
let at_0_6 = &report.threshold_sweep[1];
assert_eq!(at_0_6.min_confidence, 0.6);
assert!((at_0_6.covered_fraction - 1.0 / 3.0).abs() < 1e-9); assert!((at_0_6.abstained_fraction - 2.0 / 3.0).abs() < 1e-9);
assert_eq!(at_0_6.top1_hit_rate, Some(1.0));
assert_eq!(at_0_6.mean_reciprocal_rank, Some(1.0));
}
#[test]
fn calibration_and_threshold_sweep_are_empty_when_not_requested() {
let train = "{\"sequence_id\":\"x\",\"step\":0,\"state\":\"s\",\"action\":\"a\",\"outcome\":\"success\"}\n";
let output = evaluate(
train.as_bytes(),
train.as_bytes(),
true,
&BuildConfig::default(),
&EvalConfig::default(),
)
.unwrap();
assert!(output.report.confidence_calibration.is_empty());
assert!(output.report.threshold_sweep.is_empty());
}
#[test]
fn strict_mode_aborts_on_invalid_record_in_train_pass() {
let train = "{\"sequence_id\":\"x\",\"step\":0,\"state\":\"\",\"action\":\"a\"}\n";
let err = evaluate(
train.as_bytes(),
"".as_bytes(),
true,
&BuildConfig::default(),
&EvalConfig::default(),
)
.unwrap_err();
assert!(matches!(err, Error::EmptyState { line: 1 }));
}
#[test]
fn strict_mode_aborts_on_invalid_record_in_test_pass() {
let train = "{\"sequence_id\":\"x\",\"step\":0,\"state\":\"s\",\"action\":\"a\"}\n";
let test = "{\"sequence_id\":\"y\",\"step\":0,\"state\":\"\",\"action\":\"a\"}\n";
let err = evaluate(
train.as_bytes(),
test.as_bytes(),
true,
&BuildConfig::default(),
&EvalConfig::default(),
)
.unwrap_err();
assert!(matches!(err, Error::EmptyState { line: 1 }));
}
#[test]
fn non_strict_mode_skips_invalid_records_without_duplicating_test_pass_warnings() {
let train = "{\"sequence_id\":\"x\",\"step\":0,\"state\":\"s\",\"action\":\"a\"}\n{\"state\":\"\",\"action\":\"a\",\"sequence_id\":\"bad\",\"step\":0}\n";
let test = "{\"sequence_id\":\"y\",\"step\":0,\"state\":\"\",\"action\":\"a\"}\n";
let output = evaluate(
train.as_bytes(),
test.as_bytes(),
false,
&BuildConfig::default(),
&EvalConfig::default(),
)
.unwrap();
assert_eq!(output.warnings.len(), 1);
assert_eq!(output.warnings[0].line, 2);
}
}
#[derive(Debug, Clone)]
pub struct EvalConfig {
pub train_ratio: f64,
pub top_k: Vec<usize>,
pub calibration_bins: Option<usize>,
pub thresholds: Vec<f64>,
}
impl Default for EvalConfig {
fn default() -> Self {
Self {
train_ratio: 0.8,
top_k: vec![1, 3, 5],
calibration_bins: None,
thresholds: Vec::new(),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize)]
pub struct TopKHitRate {
pub k: usize,
pub hit_rate: Option<f64>,
}
#[derive(Debug, Clone, PartialEq, Serialize)]
pub struct CalibrationBin {
pub min_confidence: f64,
pub max_confidence: f64,
pub num_evaluated: u64,
pub top1_hit_rate: Option<f64>,
pub mean_reciprocal_rank: Option<f64>,
}
#[derive(Debug, Clone, PartialEq, Serialize)]
pub struct ThresholdSweepEntry {
pub min_confidence: f64,
pub covered_fraction: f64,
pub abstained_fraction: f64,
pub top1_hit_rate: Option<f64>,
pub mean_reciprocal_rank: Option<f64>,
}
#[derive(Debug, Clone, PartialEq, Serialize)]
pub struct EvalReport {
pub num_train_observations: u64,
pub num_test_observations: u64,
pub num_test_states: u64,
pub num_evaluated_observations: u64,
pub num_fallback_observations: u64,
pub num_test_states_with_candidates: u64,
pub coverage: Option<f64>,
pub fallback_rate: Option<f64>,
pub top1_hit_rate: Option<f64>,
pub topk_hit_rate: Vec<TopKHitRate>,
pub mean_reciprocal_rank: Option<f64>,
pub avg_rank_when_found: Option<f64>,
pub avg_confidence_on_hit: Option<f64>,
pub avg_confidence_on_miss: Option<f64>,
pub score_lift: Option<f64>,
pub success_weighted_top1_hit_rate: Option<f64>,
pub success_weighted_mean_reciprocal_rank: Option<f64>,
pub failure_agreement_top1_hit_rate: Option<f64>,
pub context_top1_hit_rate: Option<f64>,
pub context_mean_reciprocal_rank: Option<f64>,
pub hit_rate_by_matched_order: Vec<MatchedOrderHitRate>,
pub confidence_calibration: Vec<CalibrationBin>,
pub threshold_sweep: Vec<ThresholdSweepEntry>,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize)]
pub struct MatchedOrderHitRate {
pub order: usize,
pub num_evaluated: u64,
pub top1_hit_rate: Option<f64>,
}
#[derive(Debug)]
pub struct EvalOutput {
pub report: EvalReport,
pub warnings: Vec<Warning>,
}
fn is_train(sequence_id: &str, train_ratio: f64) -> bool {
let bucket = crate::hash::fnv1a(sequence_id.as_bytes()) % 100;
let train_pct = (train_ratio * 100.0).round().clamp(0.0, 100.0) as u64;
bucket < train_pct
}
#[derive(Debug, Default, Clone, Copy)]
struct CalibrationBinAcc {
num_evaluated: u64,
hit_count: u64,
reciprocal_rank_sum: f64,
}
#[derive(Debug, Default, Clone, Copy)]
struct ThresholdAcc {
covered_count: u64,
hit_count: u64,
reciprocal_rank_sum: f64,
}
#[derive(Debug, Default, Clone, Copy)]
struct MatchedOrderAcc {
num_evaluated: u64,
hit_count: u64,
}
struct EvalAccumulator<'a> {
top_k: &'a [usize],
draw_value: f64,
context_order: usize,
context_tracker: SequenceContextTracker,
context_top1_hit_count: u64,
context_reciprocal_rank_sum: f64,
matched_order_counts: HashMap<usize, MatchedOrderAcc>,
num_test_observations: u64,
test_states_seen: HashSet<String>,
states_with_candidates_count: u64,
fallback_count: u64,
evaluated_count: u64,
top1_hit_count: u64,
topk_hit_counts: HashMap<usize, u64>,
reciprocal_rank_sum: f64,
rank_sum_when_found: f64,
found_count: u64,
success_weight_sum: f64,
success_weighted_hit_sum: f64,
success_weighted_reciprocal_rank_sum: f64,
failure_count: u64,
failure_hit_count: u64,
confidence_sum_on_hit: f64,
confidence_count_on_hit: u64,
confidence_sum_on_miss: f64,
confidence_count_on_miss: u64,
score_sum_on_hit: f64,
score_count_on_hit: u64,
score_sum_on_miss: f64,
score_count_on_miss: u64,
calibration_bin_width: f64,
calibration: Vec<CalibrationBinAcc>,
thresholds: &'a [f64],
threshold_accs: Vec<ThresholdAcc>,
}
impl<'a> EvalAccumulator<'a> {
fn new(
top_k: &'a [usize],
draw_value: f64,
context_order: usize,
calibration_bins: usize,
thresholds: &'a [f64],
) -> Self {
let calibration_bin_width = if calibration_bins > 0 {
1.0 / calibration_bins as f64
} else {
0.0
};
Self {
top_k,
draw_value,
context_order,
context_tracker: SequenceContextTracker::new(context_order),
context_top1_hit_count: 0,
context_reciprocal_rank_sum: 0.0,
matched_order_counts: HashMap::new(),
num_test_observations: 0,
test_states_seen: HashSet::new(),
states_with_candidates_count: 0,
fallback_count: 0,
evaluated_count: 0,
top1_hit_count: 0,
topk_hit_counts: HashMap::new(),
reciprocal_rank_sum: 0.0,
rank_sum_when_found: 0.0,
found_count: 0,
success_weight_sum: 0.0,
success_weighted_hit_sum: 0.0,
success_weighted_reciprocal_rank_sum: 0.0,
failure_count: 0,
failure_hit_count: 0,
confidence_sum_on_hit: 0.0,
confidence_count_on_hit: 0,
confidence_sum_on_miss: 0.0,
confidence_count_on_miss: 0,
score_sum_on_hit: 0.0,
score_count_on_hit: 0,
score_sum_on_miss: 0.0,
score_count_on_miss: 0,
calibration_bin_width,
calibration: vec![CalibrationBinAcc::default(); calibration_bins],
thresholds,
threshold_accs: vec![ThresholdAcc::default(); thresholds.len()],
}
}
fn observe(&mut self, book: &PriorBook, obs: &Observation) -> Result<()> {
let window = self.context_tracker.advance(obs)?;
self.num_test_observations += 1;
let is_new_state = self.test_states_seen.insert(obs.state.clone());
let candidates = book.query(&obs.state, None);
if candidates.is_empty() {
self.fallback_count += 1;
return Ok(());
}
if is_new_state {
self.states_with_candidates_count += 1;
}
self.evaluated_count += 1;
let top1 = &candidates[0];
let is_hit = top1.action == obs.action;
if is_hit {
self.top1_hit_count += 1;
self.confidence_sum_on_hit += top1.confidence;
self.confidence_count_on_hit += 1;
if let Some(score) = obs.score {
self.score_sum_on_hit += score;
self.score_count_on_hit += 1;
}
} else {
self.confidence_sum_on_miss += top1.confidence;
self.confidence_count_on_miss += 1;
if let Some(score) = obs.score {
self.score_sum_on_miss += score;
self.score_count_on_miss += 1;
}
}
let rank = candidates
.iter()
.position(|c| c.action == obs.action)
.map(|index| index + 1);
let reciprocal_rank = rank.map_or(0.0, |r| 1.0 / r as f64);
if let Some(rank) = rank {
self.found_count += 1;
self.rank_sum_when_found += rank as f64;
self.reciprocal_rank_sum += reciprocal_rank;
for &k in self.top_k {
if rank <= k {
*self.topk_hit_counts.entry(k).or_insert(0) += 1;
}
}
}
let credit = outcome_credit(obs.outcome, self.draw_value);
self.success_weight_sum += credit;
if is_hit {
self.success_weighted_hit_sum += credit;
}
self.success_weighted_reciprocal_rank_sum += credit * reciprocal_rank;
if obs.outcome == Outcome::Failure {
self.failure_count += 1;
if is_hit {
self.failure_hit_count += 1;
}
}
if self.context_order > 0 {
let context_result = book.query_with_context(&obs.state, &window, None);
let context_top1 = &context_result.candidates[0];
let context_is_hit = context_top1.action == obs.action;
if context_is_hit {
self.context_top1_hit_count += 1;
}
let context_rank = context_result
.candidates
.iter()
.position(|c| c.action == obs.action)
.map(|index| index + 1);
self.context_reciprocal_rank_sum += context_rank.map_or(0.0, |r| 1.0 / r as f64);
let order_acc = self
.matched_order_counts
.entry(context_result.matched_order)
.or_default();
order_acc.num_evaluated += 1;
if context_is_hit {
order_acc.hit_count += 1;
}
}
if !self.calibration.is_empty() {
let bins = self.calibration.len();
let idx = ((top1.confidence / self.calibration_bin_width) as usize).min(bins - 1);
let bin = &mut self.calibration[idx];
bin.num_evaluated += 1;
if is_hit {
bin.hit_count += 1;
}
bin.reciprocal_rank_sum += reciprocal_rank;
}
for (acc, &threshold) in self.threshold_accs.iter_mut().zip(self.thresholds) {
if top1.confidence >= threshold {
acc.covered_count += 1;
if is_hit {
acc.hit_count += 1;
}
acc.reciprocal_rank_sum += reciprocal_rank;
}
}
Ok(())
}
fn finish(self, num_train_observations: u64) -> EvalReport {
let num_test_states = self.test_states_seen.len() as u64;
let coverage = ratio(
self.states_with_candidates_count as f64,
num_test_states as f64,
);
let fallback_rate = ratio(
self.fallback_count as f64,
self.num_test_observations as f64,
);
let evaluated = self.evaluated_count as f64;
let topk_hit_rate = self
.top_k
.iter()
.map(|&k| TopKHitRate {
k,
hit_rate: ratio(
*self.topk_hit_counts.get(&k).unwrap_or(&0) as f64,
evaluated,
),
})
.collect();
let score_lift = match (
ratio(self.score_sum_on_hit, self.score_count_on_hit as f64),
ratio(self.score_sum_on_miss, self.score_count_on_miss as f64),
) {
(Some(hit), Some(miss)) => Some(hit - miss),
_ => None,
};
let confidence_calibration: Vec<CalibrationBin> = self
.calibration
.iter()
.enumerate()
.map(|(i, bin)| {
let n = bin.num_evaluated as f64;
CalibrationBin {
min_confidence: i as f64 * self.calibration_bin_width,
max_confidence: (i as f64 + 1.0) * self.calibration_bin_width,
num_evaluated: bin.num_evaluated,
top1_hit_rate: ratio(bin.hit_count as f64, n),
mean_reciprocal_rank: ratio(bin.reciprocal_rank_sum, n),
}
})
.collect();
let threshold_sweep: Vec<ThresholdSweepEntry> = self
.thresholds
.iter()
.zip(self.threshold_accs.iter())
.map(|(&threshold, acc)| {
let covered_fraction =
ratio(acc.covered_count as f64, self.num_test_observations as f64)
.unwrap_or(0.0);
ThresholdSweepEntry {
min_confidence: threshold,
covered_fraction,
abstained_fraction: 1.0 - covered_fraction,
top1_hit_rate: ratio(acc.hit_count as f64, acc.covered_count as f64),
mean_reciprocal_rank: ratio(acc.reciprocal_rank_sum, acc.covered_count as f64),
}
})
.collect();
let (context_top1_hit_rate, context_mean_reciprocal_rank, hit_rate_by_matched_order) =
if self.context_order > 0 {
let mut orders: Vec<usize> = self.matched_order_counts.keys().copied().collect();
orders.sort_unstable();
let by_order = orders
.into_iter()
.map(|order| {
let acc = self.matched_order_counts[&order];
MatchedOrderHitRate {
order,
num_evaluated: acc.num_evaluated,
top1_hit_rate: ratio(acc.hit_count as f64, acc.num_evaluated as f64),
}
})
.collect();
(
ratio(self.context_top1_hit_count as f64, evaluated),
ratio(self.context_reciprocal_rank_sum, evaluated),
by_order,
)
} else {
(None, None, Vec::new())
};
EvalReport {
num_train_observations,
num_test_observations: self.num_test_observations,
num_test_states,
num_evaluated_observations: self.evaluated_count,
num_fallback_observations: self.fallback_count,
num_test_states_with_candidates: self.states_with_candidates_count,
coverage,
fallback_rate,
top1_hit_rate: ratio(self.top1_hit_count as f64, evaluated),
topk_hit_rate,
mean_reciprocal_rank: ratio(self.reciprocal_rank_sum, evaluated),
avg_rank_when_found: ratio(self.rank_sum_when_found, self.found_count as f64),
avg_confidence_on_hit: ratio(
self.confidence_sum_on_hit,
self.confidence_count_on_hit as f64,
),
avg_confidence_on_miss: ratio(
self.confidence_sum_on_miss,
self.confidence_count_on_miss as f64,
),
score_lift,
success_weighted_top1_hit_rate: ratio(
self.success_weighted_hit_sum,
self.success_weight_sum,
),
success_weighted_mean_reciprocal_rank: ratio(
self.success_weighted_reciprocal_rank_sum,
self.success_weight_sum,
),
failure_agreement_top1_hit_rate: ratio(
self.failure_hit_count as f64,
self.failure_count as f64,
),
context_top1_hit_rate,
context_mean_reciprocal_rank,
hit_rate_by_matched_order,
confidence_calibration,
threshold_sweep,
}
}
}
pub fn evaluate(
train_reader: impl Read,
test_reader: impl Read,
strict: bool,
build_config: &BuildConfig,
eval_config: &EvalConfig,
) -> Result<EvalOutput> {
let mut acc = PriorAccumulator::new(build_config)?;
let mut warnings = Vec::new();
let mut num_train_observations = 0u64;
for (index, line) in BufReader::new(train_reader).lines().enumerate() {
let line = line?;
if line.trim().is_empty() {
continue;
}
let line_no = index + 1;
match parse_line(&line, line_no) {
Ok(observation) => {
if is_train(&observation.sequence_id, eval_config.train_ratio) {
acc.observe(&observation)?;
num_train_observations += 1;
}
}
Err(err) if strict => return Err(err),
Err(err) => warnings.push(Warning {
line: line_no,
message: err.to_string(),
}),
}
}
let book = acc.finish();
let mut eval_acc = EvalAccumulator::new(
&eval_config.top_k,
build_config.draw_value,
build_config.context_order,
eval_config.calibration_bins.unwrap_or(0),
&eval_config.thresholds,
);
for (index, line) in BufReader::new(test_reader).lines().enumerate() {
let line = line?;
if line.trim().is_empty() {
continue;
}
let line_no = index + 1;
match parse_line(&line, line_no) {
Ok(observation) => {
if !is_train(&observation.sequence_id, eval_config.train_ratio) {
eval_acc.observe(&book, &observation)?;
}
}
Err(err) if strict => return Err(err),
Err(_) => {}
}
}
Ok(EvalOutput {
report: eval_acc.finish(num_train_observations),
warnings,
})
}