-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathget_simialr_entity_on_graph.py
More file actions
72 lines (60 loc) · 2.75 KB
/
Copy pathget_simialr_entity_on_graph.py
File metadata and controls
72 lines (60 loc) · 2.75 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
import os
import numpy as np
from typing import List, Set, Dict
from JudgeAgent import *
from JudgeAgent.embedding import EmbeddingClient
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--data", type=str, default="MedQA")
args = parser.parse_args()
data_name: str = args.data
# load data
data_dir = os.path.join("processed_data", data_name)
graph = Graph(data_dir)
question_path = os.path.join(data_dir, "question_with_entities.json")
save_path = os.path.join(data_dir, "questions_for_eval.json")
question_with_entities: List[Dict] = load_json(question_path)
client = EmbeddingClient(save_dir=os.path.join(data_dir, "embeddings"))
print("# Finish Load Graph and Embeddings")
# find the most similar entities on graph
entity_set: Set[str] = set()
if data_name.lower() == "quality":
for qdata in question_with_entities:
for q in qdata["questions"]:
entity_set.update([e["name"] for e in q["entities"]])
else:
for qdata in question_with_entities:
entity_set.update([e["name"] for e in qdata["entities"]])
entities: List[str] = list(entity_set)
entity_embeddings = client.get_embeddings(entities)
nodes = graph.get_node_list()
node_embeddings = client.get_embeddings(nodes)
sims: np.ndarray = np.dot(entity_embeddings, node_embeddings.T)
most_similar_node_ids = np.argmax(sims, axis=-1)
most_similar_entity_on_grpah: Dict[str, str] = {}
for ent, nid in zip(entities, most_similar_node_ids):
most_similar_entity_on_grpah[ent] = nodes[nid]
print("# Finish align entities on graph.")
# save entities
question_with_entities_on_graph: List[Dict] = []
if data_name.lower() == "quality":
for qdata in question_with_entities:
new_questions = []
for q in qdata["questions"]:
entities_on_graph = []
for e in q["entities"]:
node = graph[most_similar_entity_on_grpah[e["name"]]]
entities_on_graph.append({"name": node.name, "type": node.type})
q["entities"] = entities_on_graph
new_questions.append(q)
question_with_entities_on_graph.append({"questions": new_questions, "article": qdata["article"]})
else:
for qdata in question_with_entities:
entities_on_graph = []
for e in qdata["entities"]:
node = graph[most_similar_entity_on_grpah[e["name"]]]
entities_on_graph.append({"name": node.name, "type": node.type})
qdata["entities"] = entities_on_graph
question_with_entities_on_graph.append(qdata)
dump_json(question_with_entities_on_graph, save_path)