pdf_rag / src /retrieval.py
LightRT's picture
Changes in main.py
3301ece
Raw History Blame Contribute Delete
2.3 kB
import os
from dotenv import load_dotenv
from langchain_qdrant import QdrantVectorStore, RetrievalMode, FastEmbedSparse
from langchain_huggingface import HuggingFaceEmbeddings
from fastembed.rerank.cross_encoder import TextCrossEncoder
from qdrant_client import models
load_dotenv()
qdrant_api_key = os.getenv("QDRANT_API_KEY")
qdrant_url = os.getenv("QDRANT_URL")
class Retriever:
def __init__(self, collection_name: str = "pdf_rag"):
self.collection_name = collection_name
dense_embeddings = HuggingFaceEmbeddings(model_name="sentence-transformers/all-MiniLM-L6-v2")
sparse_embeddings = FastEmbedSparse(model_name="Qdrant/bm25")
self.vector_store = QdrantVectorStore.from_existing_collection(
embedding=dense_embeddings,
sparse_embedding=sparse_embeddings,
collection_name=collection_name,
url=qdrant_url,
api_key=qdrant_api_key,
retrieval_mode=RetrievalMode.HYBRID,
vector_name="dense",
sparse_vector_name="sparse",
)
self.reranker = TextCrossEncoder(model_name="Xenova/ms-marco-MiniLM-L-6-v2")
def retrieve(self, query: str, user_id: str, top_k: int = 5):
user_filter = models.Filter(
must=[
models.FieldCondition(
key="metadata.user_id",
match=models.MatchValue(value=user_id),
)
]
)
results = self.vector_store.similarity_search_with_score(
query,
k=20,
filter=user_filter,
)
texts = [doc.page_content for doc, _ in results]
rerank_scores = list(self.reranker.rerank(query, texts))
reranked_results = [
{
"text": doc.page_content,
"source": doc.metadata.get("source"),
"pages": doc.metadata.get("pages"),
"section": doc.metadata.get("section"),
"original_qdrant_score": score,
"rerank_score": float(rerank_score),
}
for (doc, score), rerank_score in zip(results, rerank_scores)
]
reranked_results.sort(key=lambda x: x["rerank_score"], reverse=True)
return reranked_results[:top_k]