feat: M8 complete - accuracy metrics, index tuning, gate validation
This commit is contained in:
@@ -0,0 +1,235 @@
|
||||
//! 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);
|
||||
}
|
||||
}
|
||||
@@ -12,6 +12,7 @@ pub mod gateway_queue_adapter;
|
||||
pub mod queue_worker;
|
||||
pub mod query_optimizer;
|
||||
pub mod simple_hybrid_search;
|
||||
pub mod accuracy_metrics;
|
||||
pub mod verify;
|
||||
|
||||
pub use endpoints::{IngestQueue, IngestRequest, JobStatus};
|
||||
|
||||
Reference in New Issue
Block a user