-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmetrics.py
More file actions
117 lines (95 loc) · 2.89 KB
/
Copy pathmetrics.py
File metadata and controls
117 lines (95 loc) · 2.89 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
"""Retrieval quality metrics."""
import math
from dataclasses import dataclass
@dataclass
class EvaluationResult:
"""Result of a single query evaluation."""
query: str
retrieved_docs: list[str]
relevant_docs: set[str]
recall_at_1: float
recall_at_5: float
recall_at_10: float
mrr: float
ndcg_at_10: float
@dataclass
class AggregateMetrics:
"""Aggregated metrics across multiple queries."""
num_queries: int
recall_at_1: float
recall_at_5: float
recall_at_10: float
mrr: float
ndcg_at_10: float
def format(self) -> str:
"""Format as table string."""
return f"""
# Retrieval Evaluation Results
| Metric | Value |
|--------|-------|
| Queries | {self.num_queries} |
| Recall@1 | {self.recall_at_1:.3f} |
| Recall@5 | {self.recall_at_5:.3f} |
| Recall@10 | {self.recall_at_10:.3f} |
| MRR | {self.mrr:.3f} |
| nDCG@10 | {self.ndcg_at_10:.3f} |
"""
class RecallAtK:
"""Recall@K: fraction of relevant docs found in top K results."""
@staticmethod
def compute(
retrieved: list[str],
relevant: set[str],
k: int
) -> float:
"""Compute recall at K."""
retrieved_at_k = set(retrieved[:k])
if not relevant:
return 0.0
return len(retrieved_at_k & relevant) / len(relevant)
class MRR:
"""Mean Reciprocal Rank: average of 1/rank for first relevant result."""
@staticmethod
def compute(retrieved: list[str], relevant: set[str]) -> float:
"""Compute MRR."""
if not relevant:
return 0.0
for i, doc in enumerate(retrieved, start=1):
if doc in relevant:
return 1.0 / i
return 0.0
class NDCG:
"""Normalized Discounted Cumulative Gain."""
@staticmethod
def compute(
retrieved: list[str],
relevant: set[str],
k: int = 10
) -> float:
"""Compute nDCG@K.
Assumes binary relevance (1 if relevant, 0 otherwise).
"""
# DCG: sum of relevance / log2(rank + 1)
dcg = 0.0
for i, doc in enumerate(retrieved[:k], start=1):
relevance = 1.0 if doc in relevant else 0.0
dcg += relevance / math.log2(i + 1)
# IDCG: ideal DCG (all relevant docs at top)
idcg = 0.0
for i in range(1, min(k, len(relevant)) + 1):
idcg += 1.0 / math.log2(i + 1)
if idcg == 0:
return 0.0
return dcg / idcg
def compute_metrics(
retrieved: list[str],
relevant: set[str],
) -> dict:
"""Compute all metrics for a single query."""
return {
"recall_at_1": RecallAtK.compute(retrieved, relevant, 1),
"recall_at_5": RecallAtK.compute(retrieved, relevant, 5),
"recall_at_10": RecallAtK.compute(retrieved, relevant, 10),
"mrr": MRR.compute(retrieved, relevant),
"ndcg_at_10": NDCG.compute(retrieved, relevant, 10),
}