Files
rag-agent-task/rag_agent.py
T
2026-05-25 10:57:21 +00:00

91 lines
2.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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 topk relevant documents for the query.
5. Generate an answer using OpenAI GPT3.5Turbo or GPT4 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()