91 lines
2.9 KiB
Python
91 lines
2.9 KiB
Python
"""
|
||
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()
|