feat: M8 complete - accuracy metrics, index tuning, gate validation

This commit is contained in:
2026-08-28 13:34:28 -07:00
parent f6eaae0966
commit f936931128
5 changed files with 817 additions and 0 deletions
+235
View File
@@ -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);
}
}
+1
View File
@@ -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};