Skip to content

Commit b807404

Browse files
fix: average evaluation metrics across queries and use sequential iteration for MRR (#11)
1 parent 19adc7c commit b807404

2 files changed

Lines changed: 188 additions & 44 deletions

File tree

Lines changed: 187 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,12 @@
11
use std::path::PathBuf;
22

3-
use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator};
3+
use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
44

5-
use crate::{index::InvertedIndex, query::Query, ranking::RankingAlgo};
5+
use crate::{
6+
index::InvertedIndex,
7+
query::Query,
8+
ranking::{RankingAlgo, Score},
9+
};
610

711
pub struct TestSet {
812
pub ranking_algorithm: RankingAlgo,
@@ -33,9 +37,66 @@ impl Default for Evaluation {
3337
}
3438
}
3539

40+
pub fn evaluate_query(ranked_docs: &[Score], relevant_docs: &[PathBuf]) -> Evaluation {
41+
let true_positives = ranked_docs
42+
.iter()
43+
.filter(|doc| relevant_docs.contains(&doc.doc_path))
44+
.count();
45+
46+
let false_positives = ranked_docs.len() - true_positives;
47+
let false_negatives = relevant_docs.len() - true_positives;
48+
49+
let precision = true_positives as f64 / (true_positives + false_positives) as f64;
50+
let recall = true_positives as f64 / (true_positives + false_negatives) as f64;
51+
52+
let f1_score = if precision + recall > 0.0 {
53+
2.0 * (precision * recall) / (precision + recall)
54+
} else {
55+
0.0
56+
};
57+
58+
let mrr = ranked_docs
59+
.iter()
60+
.enumerate()
61+
.find(|(_, doc)| relevant_docs.contains(&doc.doc_path))
62+
.map(|(i, _)| 1.0 / (i + 1) as f64)
63+
.unwrap_or(0.0);
64+
65+
Evaluation {
66+
precision,
67+
recall,
68+
f1_score,
69+
mrr,
70+
}
71+
}
72+
73+
pub fn average_evaluations(evaluations: &[Evaluation]) -> Evaluation {
74+
if evaluations.is_empty() {
75+
return Evaluation::default();
76+
}
77+
78+
let n = evaluations.len() as f64;
79+
let sum = evaluations
80+
.iter()
81+
.fold(Evaluation::default(), |acc, eval| Evaluation {
82+
precision: acc.precision + eval.precision,
83+
recall: acc.recall + eval.recall,
84+
f1_score: acc.f1_score + eval.f1_score,
85+
mrr: acc.mrr + eval.mrr,
86+
});
87+
88+
Evaluation {
89+
precision: sum.precision / n,
90+
recall: sum.recall / n,
91+
f1_score: sum.f1_score / n,
92+
mrr: sum.mrr / n,
93+
}
94+
}
95+
3696
impl TestSet {
3797
pub fn evaluate(&self, inverted_index: &InvertedIndex, top_n: usize) -> Evaluation {
38-
self.queries
98+
let evaluations: Vec<Evaluation> = self
99+
.queries
39100
.par_iter()
40101
.map(|query| {
41102
let ranked_docs =
@@ -47,47 +108,130 @@ impl TestSet {
47108
None => return Evaluation::default(),
48109
};
49110

50-
let true_positives = ranked_docs
51-
.par_iter()
52-
.filter(|doc| query.relevant_docs.contains(&doc.doc_path))
53-
.count();
54-
55-
let false_positives = ranked_docs.len() - true_positives;
56-
57-
let false_negatives = query.relevant_docs.len() - true_positives;
58-
59-
let precision = true_positives as f64 / (true_positives + false_positives) as f64;
60-
61-
let recall = true_positives as f64 / (true_positives + false_negatives) as f64;
62-
63-
let f1_score = if precision + recall > 0.0 {
64-
2.0 * (precision * recall) / (precision + recall)
65-
} else {
66-
0.0
67-
};
68-
69-
let mrr = ranked_docs
70-
.par_iter()
71-
.enumerate()
72-
.filter(|(_, doc)| query.relevant_docs.contains(&doc.doc_path))
73-
.map(|(i, _)| 1.0 / (i + 1) as f64)
74-
.collect::<Vec<_>>()
75-
.first()
76-
.cloned()
77-
.unwrap_or(0.0);
78-
79-
Evaluation {
80-
precision,
81-
recall,
82-
f1_score,
83-
mrr,
84-
}
111+
evaluate_query(&ranked_docs, &query.relevant_docs)
85112
})
86-
.reduce(Evaluation::default, |acc, evaluation| Evaluation {
87-
precision: acc.precision + evaluation.precision,
88-
recall: acc.recall + evaluation.recall,
89-
f1_score: acc.f1_score + evaluation.f1_score,
90-
mrr: acc.mrr + evaluation.mrr,
113+
.collect();
114+
115+
average_evaluations(&evaluations)
116+
}
117+
}
118+
119+
#[cfg(test)]
120+
mod tests {
121+
use std::path::PathBuf;
122+
123+
use super::{Evaluation, average_evaluations, evaluate_query};
124+
use crate::ranking::Score;
125+
126+
fn scored(paths: &[&str]) -> Vec<Score> {
127+
paths
128+
.iter()
129+
.enumerate()
130+
.map(|(i, p)| Score {
131+
doc_path: PathBuf::from(p),
132+
score: (paths.len() - i) as f64,
91133
})
134+
.collect()
135+
}
136+
137+
fn relevant(paths: &[&str]) -> Vec<PathBuf> {
138+
paths.iter().map(PathBuf::from).collect()
139+
}
140+
141+
#[test]
142+
fn precision_is_fraction_of_results_that_are_relevant() {
143+
// 2 of 4 results are relevant → precision = 0.5
144+
let ranked = scored(&["a.rs", "b.rs", "c.rs", "d.rs"]);
145+
let eval = evaluate_query(&ranked, &relevant(&["a.rs", "c.rs"]));
146+
147+
assert!((eval.precision - 0.5).abs() < 1e-10);
148+
}
149+
150+
#[test]
151+
fn recall_is_fraction_of_relevant_docs_retrieved() {
152+
// 1 of 3 relevant docs retrieved → recall = 1/3
153+
let ranked = scored(&["a.rs", "x.rs"]);
154+
let eval = evaluate_query(&ranked, &relevant(&["a.rs", "b.rs", "c.rs"]));
155+
156+
assert!((eval.recall - 1.0 / 3.0).abs() < 1e-10);
157+
}
158+
159+
#[test]
160+
fn f1_is_harmonic_mean_of_precision_and_recall() {
161+
let ranked = scored(&["a.rs", "x.rs"]);
162+
let eval = evaluate_query(&ranked, &relevant(&["a.rs", "b.rs"]));
163+
164+
// precision = 1/2, recall = 1/2, f1 = 2*(0.5*0.5)/(0.5+0.5) = 0.5
165+
assert!((eval.f1_score - 0.5).abs() < 1e-10);
166+
}
167+
168+
#[test]
169+
fn mrr_uses_rank_of_first_relevant_doc() {
170+
// First relevant doc is at rank 3 (0-indexed: 2) → MRR = 1/3
171+
let ranked = scored(&["x.rs", "y.rs", "a.rs", "b.rs"]);
172+
let eval = evaluate_query(&ranked, &relevant(&["a.rs", "b.rs"]));
173+
174+
assert!((eval.mrr - 1.0 / 3.0).abs() < 1e-10);
175+
}
176+
177+
#[test]
178+
fn mrr_is_one_when_first_result_is_relevant() {
179+
let ranked = scored(&["a.rs", "x.rs"]);
180+
let eval = evaluate_query(&ranked, &relevant(&["a.rs"]));
181+
182+
assert!((eval.mrr - 1.0).abs() < 1e-10);
183+
}
184+
185+
#[test]
186+
fn mrr_is_zero_when_no_relevant_docs_in_results() {
187+
let ranked = scored(&["x.rs", "y.rs"]);
188+
let eval = evaluate_query(&ranked, &relevant(&["a.rs"]));
189+
190+
assert!((eval.mrr).abs() < 1e-10);
191+
}
192+
193+
#[test]
194+
fn average_evaluations_divides_by_query_count() {
195+
let evals = vec![
196+
Evaluation {
197+
precision: 1.0,
198+
recall: 0.5,
199+
f1_score: 0.6,
200+
mrr: 1.0,
201+
},
202+
Evaluation {
203+
precision: 0.5,
204+
recall: 1.0,
205+
f1_score: 0.4,
206+
mrr: 0.5,
207+
},
208+
];
209+
let avg = average_evaluations(&evals);
210+
211+
assert!((avg.precision - 0.75).abs() < 1e-10);
212+
assert!((avg.recall - 0.75).abs() < 1e-10);
213+
assert!((avg.f1_score - 0.5).abs() < 1e-10);
214+
assert!((avg.mrr - 0.75).abs() < 1e-10);
215+
}
216+
217+
#[test]
218+
fn average_of_empty_returns_zeros() {
219+
let avg = average_evaluations(&[]);
220+
221+
assert!((avg.precision).abs() < 1e-10);
222+
assert!((avg.recall).abs() < 1e-10);
223+
}
224+
225+
#[test]
226+
fn all_metrics_bounded_zero_to_one() {
227+
let ranked = scored(&["a.rs", "b.rs", "c.rs"]);
228+
let eval = evaluate_query(&ranked, &relevant(&["a.rs", "d.rs"]));
229+
230+
for val in [eval.precision, eval.recall, eval.f1_score, eval.mrr] {
231+
assert!(
232+
(0.0..=1.0 + 1e-10).contains(&val),
233+
"metric must be in [0, 1], got {val}"
234+
);
235+
}
92236
}
93237
}

crates/core/src/evaluation/mod.rs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,4 +2,4 @@ pub mod dataset;
22
pub mod metrics;
33

44
pub use dataset::{EvaluationData, RawEvaluationData};
5-
pub use metrics::TestSet;
5+
pub use metrics::{Evaluation, TestQuery, TestSet};

0 commit comments

Comments
 (0)