236 lines
7.0 KiB
Rust
236 lines
7.0 KiB
Rust
//! 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<f32> = 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);
|
|
}
|
|
}
|