From 48a2b70c7d57fc4681502fd7e037db39f491b738 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=9A=D0=B8=D1=80=D0=B8=D0=BB=D0=BB=20=D0=A0=D0=BE=D0=BC?= =?UTF-8?q?=D0=B0=D0=BD=D0=BE=D0=B2?= Date: Mon, 25 May 2026 10:57:21 +0000 Subject: [PATCH] Add rag_agent.py --- rag_agent.py | 90 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 90 insertions(+) create mode 100644 rag_agent.py diff --git a/rag_agent.py b/rag_agent.py new file mode 100644 index 0000000..7f5c3e9 --- /dev/null +++ b/rag_agent.py @@ -0,0 +1,90 @@ +""" +Simple Retrieval-Augmented Generation (RAG) agent. + +Usage: + python rag_agent.py --data_dir /path/to/docs --query "Your question" + +The script will: + 1. Load all .txt files from data_dir. + 2. Create embeddings using OpenAI's text-embedding-ada-002. + 3. Store them in a FAISS vector store. + 4. Retrieve top‑k relevant documents for the query. + 5. Generate an answer using OpenAI GPT‑3.5‑Turbo or GPT‑4 via LangChain. +""" + +import os +import argparse +from pathlib import Path + +from langchain.document_loaders import TextLoader +from langchain.embeddings import OpenAIEmbeddings +from langchain.vectorstores import FAISS +from langchain.llms import OpenAI +from langchain.chains import RetrievalQA +from langchain.prompts import PromptTemplate + + +def load_documents(data_dir: str): + """Load all .txt files from data_dir. + + Args: + data_dir: Directory containing text files. + + Returns: + List of langchain Document objects. + """ + docs = [] + for file_path in Path(data_dir).rglob("*.txt"): + loader = TextLoader(str(file_path)) + docs.extend(loader.load()) + return docs + + +def main(): + parser = argparse.ArgumentParser(description="Simple RAG agent using LangChain and OpenAI") + parser.add_argument("--data_dir", required=True, help="Directory with .txt documents") + parser.add_argument("--query", required=True, help="Query to answer") + parser.add_argument("--model", default="gpt-3.5-turbo", help="OpenAI model for generation (gpt-3.5-turbo or gpt-4)") + parser.add_argument("--k", type=int, default=5, help="Number of documents to retrieve") + args = parser.parse_args() + + if "OPENAI_API_KEY" not in os.environ: + raise ValueError("OPENAI_API_KEY environment variable not set") + + # 1. Load documents + print("Loading documents...") + docs = load_documents(args.data_dir) + if not docs: + raise ValueError("No documents found in the specified directory") + + # 2. Create embeddings and vector store + print("Creating embeddings and vector store...") + embeddings = OpenAIEmbeddings() + vector_store = FAISS.from_documents(docs, embeddings) + + # 3. Build Retriever and QA chain + retriever = vector_store.as_retriever(search_kwargs={"k": args.k}) + prompt_template = PromptTemplate( + input_variables=["context", "question"], + template=""" +You are an AI assistant. Use the following context to answer the question. + +Context:\n{context}\n\nQuestion: {question}\n\nAnswer:\n""" + ) + llm = OpenAI(model_name=args.model, temperature=0) + qa_chain = RetrievalQA.from_chain_type( + llm=llm, + chain_type="stuff", + retriever=retriever, + return_source_documents=False, + chain_type_kwargs={"prompt": prompt_template}, + ) + + # 4. Run query + print("Running query...") + result = qa_chain(args.query) + print("\nAnswer:\n", result["answer"]) + + +if __name__ == "__main__": + main()