Kelvin-programmer/rag-chatbot
0
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 