CoolFace
Apppublic

Kelvin-programmer/rag-chatbot

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
rag_engine.py94 linesDownload Raw Back to src
1"""Core RAG pipeline: retrieval + generation."""2 3import logging4 5import torch6from transformers import AutoModelForSeq2SeqLM, AutoTokenizer7 8from .config import Settings9from .pdf_processor import PDFProcessor10from .vector_store import VectorStore11 12logger = logging.getLogger(__name__)13 14PROMPT_TEMPLATE = """Answer the question based only on the provided context. \15Be concise and accurate. If the context does not contain enough information, \16say "I don't have enough information to answer that question."17 18Context:19{context}20 21Question: {question}22 23Answer:"""24 25 26class RAGEngine:27    """Orchestrates the retrieval-augmented generation pipeline."""28 29    def __init__(self, settings: Settings):30        self.settings = settings31 32        logger.info("Loading embedding model: %s", settings.embedding_model)33        self.vector_store = VectorStore(34            model_name=settings.embedding_model,35            persist_dir=settings.vector_store_path,36        )37 38        self.pdf_processor = PDFProcessor(39            chunk_size=settings.chunk_size,40            chunk_overlap=settings.chunk_overlap,41        )42 43        logger.info("Loading LLM: %s", settings.llm_model)44        self.tokenizer = AutoTokenizer.from_pretrained(settings.llm_model)45        self.model = AutoModelForSeq2SeqLM.from_pretrained(settings.llm_model)46        self.device = "cuda" if torch.cuda.is_available() else "cpu"47        self.model.to(self.device)48 49        logger.info("RAG Engine ready (device=%s)", self.device)50 51    def ingest_pdf(self, pdf_path: str) -> int:52        """Parse a PDF, chunk its text, and add to the vector store."""53        results = self.pdf_processor.extract_chunks(pdf_path)54        texts = [r["text"] for r in results]55        metadata = [r["metadata"] for r in results]56        count = self.vector_store.add_documents(texts, metadata)57        self.vector_store.save()58        return count59 60    def query(self, question: str) -> dict:61        """Retrieve relevant context and generate a grounded answer."""62        results = self.vector_store.search(question, top_k=self.settings.top_k)63 64        if not results:65            return {66                "answer": "No documents have been loaded. Please upload a PDF first.",67                "sources": [],68            }69 70        context = "\n\n".join(r["text"] for r in results)71        prompt = PROMPT_TEMPLATE.format(context=context, question=question)72 73        inputs = self.tokenizer(74            prompt, return_tensors="pt", max_length=512, truncation=True75        ).to(self.device)76 77        with torch.no_grad():78            outputs = self.model.generate(79                **inputs, max_new_tokens=self.settings.max_tokens80            )81 82        answer = self.tokenizer.decode(outputs[0], skip_special_tokens=True).strip()83 84        sources = [85            {86                "text": r["text"][:200] + ("..." if len(r["text"]) > 200 else ""),87                "score": r["score"],88                "metadata": r["metadata"],89            }90            for r in results91        ]92 93        return {"answer": answer, "sources": sources}94