ContiAI / rag /rag_engine.py
ziadsameh32's picture
Add login page
325b94c
Raw
History Blame Contribute Delete
3.54 kB
import os
import uuid
from dataclasses import asdict
from typing import List, Optional
from sentence_transformers import SentenceTransformer
from qdrant_client import QdrantClient
from qdrant_client.http import models as qm
from schemas import ChunkRecord
from .ocr import mistral_ocr_pdf
from .preprocess import normalize_arabic, drop_common_headers_footers
from .chuncking import chunk_pages
class ArabicBookRAG:
def __init__(
self,
collection_name: str, # ๐Ÿ‘ˆ ู…ู‡ู…
embedding_model="sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2",
):
self.collection = collection_name
self.embedder = SentenceTransformer(embedding_model)
self.qdrant = QdrantClient(
url=os.environ["QDRANT_URL"],
api_key=os.environ["QDRANT_API_KEY"],
)
self._ensure_collection()
def _ensure_collection(self):
dim = self.embedder.get_sentence_embedding_dimension()
collections = [c.name for c in self.qdrant.get_collections().collections]
if self.collection not in collections:
self.qdrant.create_collection(
collection_name=self.collection,
vectors_config=qm.VectorParams(size=dim, distance=qm.Distance.COSINE),
)
# =========================
# Ingest PDF (In-Memory)
# =========================
def ingest_pdf(self, pdf_bytes: bytes, meta):
pages = mistral_ocr_pdf(pdf_bytes)
pages = [normalize_arabic(p) for p in pages]
pages = drop_common_headers_footers(pages)
chunks = chunk_pages(pages)
records = [
ChunkRecord(
chunk_id=str(uuid.uuid4()),
doc_id=meta.doc_id,
page_start=ps,
page_end=pe,
text=txt,
title=meta.title,
author=meta.author,
year=meta.year,
)
for txt, ps, pe in chunks
]
vectors = self.embedder.encode(
[r.text for r in records], normalize_embeddings=True
)
self.qdrant.upsert(
collection_name=self.collection,
points=[
qm.PointStruct(
id=r.chunk_id,
vector=v.tolist(),
payload=asdict(r),
)
for r, v in zip(records, vectors)
],
)
return {"pages": len(pages), "chunks": len(records)}
# =========================
# Retrieve
# =========================
def retrieve(
self, queries: List[str], doc_id: Optional[str] = None, top_k: int = 8
):
must = []
if doc_id:
must.append(
qm.FieldCondition(key="doc_id", match=qm.MatchValue(value=doc_id))
)
flt = qm.Filter(must=must) if must else None
hits = []
for q in queries:
qn = normalize_arabic(q)
vec = self.embedder.encode([qn], normalize_embeddings=True)[0]
res = self.qdrant.query_points(
collection_name=self.collection,
query=vec.tolist(),
limit=top_k,
with_payload=True,
query_filter=flt,
).points
hits.extend(res)
return hits
# =========================
# Delete whole book
# =========================
def delete_collection(self):
self.qdrant.delete_collection(self.collection)