Spaces:
Sleeping
Sleeping
making our own reranker for faster result
Browse files
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
|
| 22 |
# from langchain_chroma import Chroma
|
| 23 |
from langchain_community.retrievers import BM25Retriever
|
| 24 |
-
from langchain.retrievers import EnsembleRetriever
|
| 25 |
-
from
|
| 26 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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 |
-
#
|
| 184 |
-
#
|
| 185 |
-
#
|
| 186 |
-
#
|
| 187 |
-
#
|
| 188 |
-
#
|
| 189 |
-
|
| 190 |
-
#
|
| 191 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|