11use 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
711pub 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+
3696impl 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}
0 commit comments