-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
119 lines (94 loc) · 4.24 KB
/
Copy pathmain.py
File metadata and controls
119 lines (94 loc) · 4.24 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
118
119
from dotenv import load_dotenv
load_dotenv()
import os
import logging
import time
import json
import torch
from datetime import datetime
from pathlib import Path
from sentence_transformers import SentenceTransformer
from src.retriever import NVEmbedEncoder, QwenEncoder
from src.BrowseNet import BrowseNet
from src.NaiveRAG import NaiveRAG
ROOT_DIR = Path(__file__).resolve().parent
if __name__ == '__main__':
dataset = os.environ['DATASET']
alpha = float(os.environ['ALPHA'])
ner_model = os.environ['NER_MODEL']
sem_model = os.environ['SEM_MODEL']
n_subgraphs = int(os.environ['N_SUBGRAPHS'])
subquery_model = os.environ['SUBQUERY_MODEL']
colbert_threshold = float(os.environ['COLBERT_THRESHOLD'])
retrieval_method = os.environ['RETRIEVAL_METHOD']
n_chunks = int(os.environ['N_CHUNKS'])
llm = os.environ['LLM']
model_name = os.environ['MODEL']
os.makedirs(f'logs/{dataset}', exist_ok=True)
os.makedirs(f'results/{dataset}', exist_ok=True)
os.makedirs(f'artifacts/{dataset}', exist_ok=True)
logging.basicConfig(level=logging.INFO, filename='logs/{}/{}.log'.format(dataset,datetime.now().timestamp()),
filemode='a', format='%(asctime)s - %(name)s - %(levelname)s - %(message)s')
logging.getLogger("httpx").setLevel(logging.WARNING)
logger = logging.getLogger(__name__)
device = os.getenv('DEVICE')
device = device if 'cuda' in device and torch.cuda.is_available() else 'cpu'
logger.info(f"::> PIPELINE INFO:\n\
Dataset: {dataset}\n\
Retrieval Method: {retrieval_method}\n\
Sem Model: {sem_model}\n\
Device: {device}\n\
NER Model: {ner_model}\n\
Subquery Model: {subquery_model}\n\
Colbert Threshold: {colbert_threshold}\n\
Number of Subgraphs: {n_subgraphs}\n\
Alpha: {alpha}\n\
Number of Chunks: {n_chunks}\n\
")
start_time = time.time()
logger.info(f"::> Time taken to load the encoder: {time.time()-start_time} seconds")
if retrieval_method.lower() == 'browsenet':
logger.info(f"::> Initializing BrowseNet...")
browsenet = BrowseNet(
dataset = dataset,
device=device,
ner_model = ner_model,
sem_model= sem_model,
subquery_model = subquery_model,
colbert_threshold = colbert_threshold,
n_subgraphs = n_subgraphs,
alpha = alpha
)
logger.info(f"::> Indexing...")
browsenet.index()
questions = json.load(open(ROOT_DIR / 'datasets' / dataset / 'questions.json','r'))
logger.info(f"::> Total Questions: {len(questions)}")
logger.info(f"::> Starting Retrieval...")
split_queries, retrieved_corpus = browsenet.retrieve(questions)
logger.info(f"::> Starting Retrieval Evaluation...")
result_dict = browsenet.retrieval_eval(questions, retrieved_corpus)
logger.info(f"::> Starting QA...")
browsenet.qa(questions, split_queries, retrieved_corpus, n_chunks, llm, model_name)
logger.info(f"::> Starting QA Evaluation...")
em, f1 = browsenet.qa_eval(dataset, n_chunks)
else:
logger.info(f"::> Initializing NaiveRAG...")
naiverag = NaiveRAG(
dataset = dataset,
device=device,
sem_model= sem_model,
alpha = alpha
)
logger.info(f"::> Indexing...")
naiverag.index()
questions = json.load(open(ROOT_DIR / 'datasets' / dataset / 'questions.json','r'))
logger.info(f"::> Total Questions: {len(questions)}")
logger.info(f"::> Starting Retrieval...")
split_queries,retrieved_corpus = naiverag.retrieve(questions)
logger.info(f"::> Starting Retrieval Evaluation...")
result_dict = naiverag.retrieval_eval(questions, retrieved_corpus)
logger.info(f"::> Starting QA...")
naiverag.qa(questions, split_queries, retrieved_corpus, n_chunks, llm, model_name)
logger.info(f"::> Starting QA Evaluation...")
em, f1 = naiverag.qa_eval(dataset, n_chunks)
logger.info(f"::> Total time taken for the pipeline: {time.time() - start_time} seconds.\n\n\n")