singhankur01 commited on
Commit
5225cff
·
verified ·
1 Parent(s): 45c44eb

making our own reranker for faster result

Browse files
Files changed (1) hide show
  1. app.py +51 -26
app.py CHANGED
@@ -11,19 +11,23 @@ from operator import itemgetter
11
  from fastapi import FastAPI, Depends, HTTPException, Header
12
  from fastapi.responses import JSONResponse
13
 
14
- # Make sure you have these files in a 'utils' folder
15
  from utils.DocsLoader import load_and_chunk
16
  from utils.Schemas import RunRequest, RunResponse
17
  # from concurrent.futures import ThreadPoolExecutor
18
  from langchain_community.vectorstores import FAISS
19
  from langchain.schema import Document
20
  from langchain_google_genai import ChatGoogleGenerativeAI
21
- from langchain_huggingface import HuggingFaceEmbeddings # Correct new import
22
  # from langchain_chroma import Chroma
23
  from langchain_community.retrievers import BM25Retriever
24
- from langchain.retrievers import EnsembleRetriever, ContextualCompressionRetriever
25
- from langchain.retrievers.document_compressors import CrossEncoderReranker
26
- from langchain_community.cross_encoders import HuggingFaceCrossEncoder
 
 
 
 
 
27
  from langchain.prompts import PromptTemplate
28
 
29
 
@@ -48,9 +52,10 @@ async def lifespan(app: FastAPI):
48
 
49
  # Load models into the shared dictionary
50
  ml_models["embedder"] = HuggingFaceEmbeddings(model_name="BAAI/bge-base-en-v1.5")
51
- cross_encoder_model = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-base")
 
52
  # cross_encoder_model = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-large")
53
- ml_models["reranker_compressor"] = CrossEncoderReranker(model=cross_encoder_model, top_n=5)
54
  ml_models["llm"] = ChatGoogleGenerativeAI(model="gemini-2.0-flash", api_key=GOOGLE_API_KEY)
55
  ml_models["prompt_template"] = PromptTemplate.from_template(
56
  """
@@ -165,30 +170,50 @@ async def run_hackrx(req: RunRequest):
165
  keyword_retriever.k = 5
166
  # dense_retriever = Chroma.from_documents(documents=chunks, embedding=ml_models["embedder"]).as_retriever()
167
  ensemble_retriever = EnsembleRetriever(retrievers=[keyword_retriever, dense_retriever], weights=[0.4, 0.6])
168
-
169
- compression_retriever = ContextualCompressionRetriever(
170
- base_retriever=ensemble_retriever, base_compressor=ml_models["reranker_compressor"]
171
- )
172
  # compression_retriever = ContextualCompressionRetriever(
173
  # base_retriever=ensemble_retriever, base_compressor=ml_models["reranker_compressor"]
174
  # )
175
 
176
- # Define the RAG chain using pre-loaded components
177
- hybrid_rag_chain = (
178
- {"context": itemgetter("full_query") | compression_retriever, "full_query": itemgetter("full_query")}
179
- | ml_models["prompt_template"]
180
- | ml_models["llm"]
181
- )
182
 
183
- # answers = []
184
- # for q in req.questions:
185
- # try:
186
- # result = await hybrid_rag_chain.ainvoke({"full_query": q})
187
- # parsed = parse_llm_response(result.content)
188
- # answers.append(parsed)
189
- # except Exception as e:
190
- # return JSONResponse({"error": str(e)}, status_code=500)
191
- tasks = [hybrid_rag_chain.ainvoke({"full_query": q}) for q in req.questions]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
192
  results = await asyncio.gather(*tasks)
193
 
194
  # Extract the content from each result and parse it
 
11
  from fastapi import FastAPI, Depends, HTTPException, Header
12
  from fastapi.responses import JSONResponse
13
 
 
14
  from utils.DocsLoader import load_and_chunk
15
  from utils.Schemas import RunRequest, RunResponse
16
  # from concurrent.futures import ThreadPoolExecutor
17
  from langchain_community.vectorstores import FAISS
18
  from langchain.schema import Document
19
  from langchain_google_genai import ChatGoogleGenerativeAI
20
+ from langchain_huggingface import HuggingFaceEmbeddings
21
  # from langchain_chroma import Chroma
22
  from langchain_community.retrievers import BM25Retriever
23
+ from langchain.retrievers import EnsembleRetriever
24
+ from sklearn.metrics.pairwise import cosine_similarity
25
+ import numpy as np
26
+
27
+ ### to make it faster we are now using our built reranker thats why commenting the imports below
28
+ # from langchain.retrievers import ContextualCompressionRetriever
29
+ # from langchain.retrievers.document_compressors import CrossEncoderReranker
30
+ # from langchain_community.cross_encoders import HuggingFaceCrossEncoder
31
  from langchain.prompts import PromptTemplate
32
 
33
 
 
52
 
53
  # Load models into the shared dictionary
54
  ml_models["embedder"] = HuggingFaceEmbeddings(model_name="BAAI/bge-base-en-v1.5")
55
+ ### to make it faster we are now using our built reranker thats why commenting the code below
56
+ # cross_encoder_model = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-base")
57
  # cross_encoder_model = HuggingFaceCrossEncoder(model_name="BAAI/bge-reranker-large")
58
+ # ml_models["reranker_compressor"] = CrossEncoderReranker(model=cross_encoder_model, top_n=5)
59
  ml_models["llm"] = ChatGoogleGenerativeAI(model="gemini-2.0-flash", api_key=GOOGLE_API_KEY)
60
  ml_models["prompt_template"] = PromptTemplate.from_template(
61
  """
 
170
  keyword_retriever.k = 5
171
  # dense_retriever = Chroma.from_documents(documents=chunks, embedding=ml_models["embedder"]).as_retriever()
172
  ensemble_retriever = EnsembleRetriever(retrievers=[keyword_retriever, dense_retriever], weights=[0.4, 0.6])
173
+ ### to make it faster we are now using our built reranker thats why commenting the code below
 
 
 
174
  # compression_retriever = ContextualCompressionRetriever(
175
  # base_retriever=ensemble_retriever, base_compressor=ml_models["reranker_compressor"]
176
  # )
177
 
 
 
 
 
 
 
178
 
179
+ # Define the RAG chain using pre-loaded components
180
+ # hybrid_rag_chain = (
181
+ # {"context": itemgetter("full_query") | compression_retriever, "full_query": itemgetter("full_query")}
182
+ # | ml_models["prompt_template"]
183
+ # | ml_models["llm"]
184
+ # )
185
+
186
+ ######## OUR SELF RERANKER ######################################################################
187
+ #Embed all questions at once
188
+ question_embeddings = ml_models["embedder"].embed_documents(req.questions)
189
+
190
+ # For each question, retrieve and rerank with cosine
191
+ retrieved_chunks_all = []
192
+ for i, question in enumerate(req.questions):
193
+ docs = ensemble_retriever.get_relevant_documents(question)
194
+ doc_texts = [doc.page_content for doc in docs]
195
+
196
+ doc_embeddings = ml_models["embedder"].embed_documents(doc_texts)
197
+ sims = cosine_similarity([question_embeddings[i]], doc_embeddings)[0]
198
+
199
+ top_k = 5
200
+ top_indices = np.argsort(sims)[-top_k:][::-1]
201
+ top_chunks = [doc_texts[j] for j in top_indices]
202
+
203
+ # Join for context
204
+ joined_context = "\n\n".join(top_chunks)
205
+ retrieved_chunks_all.append(joined_context)
206
+
207
+ ####################################################################################################################
208
+
209
+ tasks = []
210
+ for i in range(len(req.questions)):
211
+ prompt_input = {
212
+ "full_query": req.questions[i],
213
+ "context": retrieved_chunks_all[i]
214
+ }
215
+ tasks.append(ml_models["llm"].ainvoke(ml_models["prompt_template"].format_prompt(**prompt_input)))
216
+ # tasks = [hybrid_rag_chain.ainvoke({"full_query": q}) for q in req.questions]
217
  results = await asyncio.gather(*tasks)
218
 
219
  # Extract the content from each result and parse it