//! M8.8 — Accuracy Metrics: NDCG, MRR, Precision@K, Recall@K //! //! Measures search quality for hybrid search tuning and benchmarking. use serde::{Deserialize, Serialize}; use std::collections::HashSet; /// Accuracy metrics for search results #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AccuracyMetrics { pub query_id: String, pub ndcg_10: f32, // NDCG@10 pub mrr: f32, // Mean Reciprocal Rank pub precision_10: f32, // Precision@10 pub recall_10: f32, // Recall@10 pub relevant_count: usize, // Total relevant documents pub retrieved_count: usize, // Documents retrieved } impl Default for AccuracyMetrics { fn default() -> Self { Self { query_id: String::new(), ndcg_10: 0.0, mrr: 0.0, precision_10: 0.0, recall_10: 0.0, relevant_count: 0, retrieved_count: 0, } } } /// Calculate NDCG@K (Normalized Discounted Cumulative Gain) /// /// Measures ranking quality by penalizing misranked relevant documents. /// 1.0 = perfect ranking, 0.0 = no relevant docs in top-k pub fn ndcg_at_k(relevant_ids: &[&str], retrieved_ids: &[&str], k: usize) -> f32 { let relevant_set: HashSet<_> = relevant_ids.iter().collect(); // Calculate DCG@K let mut dcg = 0.0; for (i, doc_id) in retrieved_ids.iter().take(k).enumerate() { if relevant_set.contains(doc_id) { dcg += 1.0 / ((i as f32 + 2.0).log2()); } } // Calculate IDCG@K (ideal ranking: all relevant docs first) let mut idcg = 0.0; for i in 0..relevant_ids.len().min(k) { idcg += 1.0 / ((i as f32 + 2.0).log2()); } if idcg == 0.0 { 0.0 } else { dcg / idcg } } /// Calculate MRR (Mean Reciprocal Rank) /// /// Position of first relevant document. 1.0 if first, 0.5 if second, etc. pub fn mrr(relevant_ids: &[&str], retrieved_ids: &[&str]) -> f32 { let relevant_set: HashSet<_> = relevant_ids.iter().collect(); for (i, doc_id) in retrieved_ids.iter().enumerate() { if relevant_set.contains(doc_id) { return 1.0 / (i as f32 + 1.0); } } 0.0 } /// Calculate Precision@K /// /// Fraction of top-k results that are relevant. pub fn precision_at_k(relevant_ids: &[&str], retrieved_ids: &[&str], k: usize) -> f32 { let relevant_set: HashSet<_> = relevant_ids.iter().collect(); let mut hits = 0; for doc_id in retrieved_ids.iter().take(k) { if relevant_set.contains(doc_id) { hits += 1; } } hits as f32 / k as f32 } /// Calculate Recall@K /// /// Fraction of relevant documents found in top-k results. pub fn recall_at_k(relevant_ids: &[&str], retrieved_ids: &[&str], k: usize) -> f32 { if relevant_ids.is_empty() { return 0.0; } let relevant_set: HashSet<_> = relevant_ids.iter().collect(); let mut hits = 0; for doc_id in retrieved_ids.iter().take(k) { if relevant_set.contains(doc_id) { hits += 1; } } hits as f32 / relevant_ids.len() as f32 } /// Summary statistics across multiple queries #[derive(Debug, Clone, Serialize, Deserialize)] pub struct BenchmarkSummary { pub query_count: usize, pub mean_ndcg_10: f32, pub mean_mrr: f32, pub mean_precision_10: f32, pub mean_recall_10: f32, pub median_ndcg_10: f32, } impl BenchmarkSummary { pub fn from_metrics(metrics: &[AccuracyMetrics]) -> Self { if metrics.is_empty() { return Self { query_count: 0, mean_ndcg_10: 0.0, mean_mrr: 0.0, mean_precision_10: 0.0, mean_recall_10: 0.0, median_ndcg_10: 0.0, }; } let sum_ndcg: f32 = metrics.iter().map(|m| m.ndcg_10).sum(); let sum_mrr: f32 = metrics.iter().map(|m| m.mrr).sum(); let sum_prec: f32 = metrics.iter().map(|m| m.precision_10).sum(); let sum_rec: f32 = metrics.iter().map(|m| m.recall_10).sum(); let mut ndcg_values: Vec = metrics.iter().map(|m| m.ndcg_10).collect(); ndcg_values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); let median_ndcg = if ndcg_values.len() % 2 == 0 { (ndcg_values[ndcg_values.len() / 2 - 1] + ndcg_values[ndcg_values.len() / 2]) / 2.0 } else { ndcg_values[ndcg_values.len() / 2] }; Self { query_count: metrics.len(), mean_ndcg_10: sum_ndcg / metrics.len() as f32, mean_mrr: sum_mrr / metrics.len() as f32, mean_precision_10: sum_prec / metrics.len() as f32, mean_recall_10: sum_rec / metrics.len() as f32, median_ndcg_10: median_ndcg, } } } #[cfg(test)] mod tests { use super::*; #[test] fn test_ndcg_perfect_ranking() { let relevant = vec!["doc1", "doc2", "doc3"]; let retrieved = vec!["doc1", "doc2", "doc3", "doc4"]; let ndcg = ndcg_at_k(&relevant, &retrieved, 10); assert!((ndcg - 1.0).abs() < 0.001); } #[test] fn test_ndcg_worst_ranking() { let relevant = vec!["doc1", "doc2", "doc3"]; let retrieved = vec!["doc4", "doc5", "doc6", "doc7"]; let ndcg = ndcg_at_k(&relevant, &retrieved, 10); assert!(ndcg < 0.001); } #[test] fn test_mrr_first_position() { let relevant = vec!["doc1"]; let retrieved = vec!["doc1", "doc2"]; assert!((mrr(&relevant, &retrieved) - 1.0).abs() < 0.001); } #[test] fn test_mrr_second_position() { let relevant = vec!["doc1"]; let retrieved = vec!["doc2", "doc1"]; assert!((mrr(&relevant, &retrieved) - 0.5).abs() < 0.001); } #[test] fn test_precision_at_10() { let relevant = vec!["doc1", "doc2"]; let retrieved = vec!["doc1", "doc3", "doc4", "doc5", "doc2", "doc6"]; let prec = precision_at_k(&relevant, &retrieved, 10); assert!((prec - 0.2).abs() < 0.001); // 2/10 = 0.2 } #[test] fn test_recall_at_10() { let relevant = vec!["doc1", "doc2", "doc3"]; let retrieved = vec!["doc1", "doc4", "doc2"]; let rec = recall_at_k(&relevant, &retrieved, 10); assert!((rec - (2.0 / 3.0)).abs() < 0.001); // 2/3 = 0.667 } #[test] fn test_benchmark_summary() { let metrics = vec![ AccuracyMetrics { ndcg_10: 0.9, mrr: 1.0, precision_10: 0.8, recall_10: 0.7, ..Default::default() }, AccuracyMetrics { ndcg_10: 0.7, mrr: 0.5, precision_10: 0.6, recall_10: 0.5, ..Default::default() }, ]; let summary = BenchmarkSummary::from_metrics(&metrics); assert_eq!(summary.query_count, 2); assert!((summary.mean_ndcg_10 - 0.8).abs() < 0.001); } }