-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathevaluate3.py
More file actions
102 lines (92 loc) · 3.45 KB
/
Copy pathevaluate3.py
File metadata and controls
102 lines (92 loc) · 3.45 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
import faiss
import json
import numpy as np
import ir_datasets
from sentence_transformers import SentenceTransformer
from tqdm import tqdm
#english to japanese based off the save_english_queries.py
#the comments are for my reference to learn as i build hehe
model = SentenceTransformer("sentence-transformers/LaBSE", device="cuda")
index = faiss.read_index("index/corpus.index")
with open("index/doc_ids.json") as f:
doc_ids = json.load(f)
with open("index/english_queries.json") as f:
english_queries = json.load(f)
dataset = ir_datasets.load("mr-tydi/ja/dev")
#creating a dict for the qrels for easier access cus qrels in the dataset have multiple entries per query. now we can easily get the relevant doc ids for each query by looking up the query id in this dict.
qrels={}
for qrel in dataset.qrels_iter():
if qrel.query_id not in qrels:
qrels[qrel.query_id] = set()
qrels[qrel.query_id].add(qrel.doc_id)
#all the queries from the dataset
queries = [(q["query_id"], q["text"]) for q in english_queries]
print(f"Evaluable queries: {len(queries)}")
query_texts = []
for q in queries:
query_texts.append(q[1])
query_vecs = model.encode(query_texts, normalize_embeddings=True,device="cuda", show_progress_bar=True)
# Retrieve top 100
scores, indices = index.search(query_vecs.astype(np.float32), 100)
mrr_scores = []
recall_scores = []
for i, (qid, _) in enumerate(queries):
relevant = qrels.get(qid, set())
top100 = [doc_ids[idx] for idx in indices[i]]
#Recall@100
hits = sum(1 for d in top100 if d in relevant)
recall_scores.append(hits / len(relevant))
# MRR@10
mrr = 0.0
for rank, doc_id in enumerate(top100[:10], start=1):
if doc_id in relevant:
mrr = 1.0 / rank
break
mrr_scores.append(mrr)
print(f"MRR@10: {np.mean(mrr_scores):.4f}")
print(f"Recall@100: {np.mean(recall_scores):.4f}")
from rocchio import rocchio
# Evaluate with simulated relevance feedback
mrr_rf_scores = []
recall_rf_scores = []
for i, (qid, _) in enumerate(queries):
relevant = qrels.get(qid, set())
top100_ids = [doc_ids[idx] for idx in indices[i]]
top100_vecs = np.array([index.reconstruct(int(idx)) for idx in indices[i]])
# Simulate feedback: mark top 10 as relevant/non-relevant based on qrels
relevant_vecs = []
non_relevant_vecs = []
for j, doc_id in enumerate(top100_ids[:10]):
if doc_id in relevant:
relevant_vecs.append(top100_vecs[j])
else:
non_relevant_vecs.append(top100_vecs[j])
# If no relevant docs in top 10, skip feedback
if not relevant_vecs:
mrr_rf_scores.append(mrr_scores[i])
recall_rf_scores.append(recall_scores[i])
continue
# Apply Rocchio
new_query_vec = rocchio(
query_vecs[i],
np.array(relevant_vecs),
np.array(non_relevant_vecs)
)
# Re-retrieve
new_scores, new_indices = index.search(
new_query_vec.reshape(1, -1).astype(np.float32), 100
)
new_top100 = [doc_ids[idx] for idx in new_indices[0]]
# Recall@100
hits = sum(1 for d in new_top100 if d in relevant)
recall_rf_scores.append(hits / len(relevant))
# MRR@10
mrr = 0.0
for rank, doc_id in enumerate(new_top100[:10], start=1):
if doc_id in relevant:
mrr = 1.0 / rank
break
mrr_rf_scores.append(mrr)
print(f"\n After Relevance Feedback")
print(f"MRR@10: {np.mean(mrr_rf_scores):.4f}")
print(f"Recall@100: {np.mean(recall_rf_scores):.4f}")