llm-rag / app.py
Agboola
update
8a19df4
Raw
History Blame Contribute Delete
4.4 kB
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from langchain_community.document_loaders import PyPDFLoader
from langchain_text_splitters import RecursiveCharacterTextSplitter
# from langchain_community.embeddings import HuggingFaceEmbeddings
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_community.vectorstores import Chroma
from langchain_core.prompts import PromptTemplate
from langchain_core.runnables import RunnablePassthrough
from langchain_core.output_parsers import StrOutputParser
from transformers import AutoTokenizer, AutoModelForCausalLM, pipeline
from langchain_huggingface import HuggingFacePipeline
import torch
from transformers import StoppingCriteria, StoppingCriteriaList
from langchain_core.runnables import Runnable
import traceback
app = FastAPI()
class StopOnPunctuationOrMaxTokens(StoppingCriteria):
def __init__(self, max_tokens: int = 256):
self.max_tokens = max_tokens
def __call__(self, input_ids, scores, **kwargs):
last_token = input_ids[0][-1]
return last_token == 46 or last_token == 63 # IDs '.' and '?'
# # Stop if the last token is a period (.), exclamation mark (!), or question mark (?)
# return last_token == 46 or last_token == 33 or last_token == 63 # Period, '!', or '?'
class CleanStrOutputParser(Runnable):
def invoke(self, input_text: str, config=None) -> str:
# Remove quotes and extra whitespace
text = input_text.strip()
if text.startswith('"') and text.endswith('"'):
text = text[1:-1]
if "." in text:
text = text[: text.rfind(".") + 1]
return text
# -----------------------
# Schema Validation
# -----------------------
class Query(BaseModel):
question: str
# -----------------------
# Load documents
# -----------------------
loader = PyPDFLoader("seplat.pdf")
docs = loader.load()
text_splitter = RecursiveCharacterTextSplitter(
chunk_size=1000,
chunk_overlap=200
)
doc_chunks = text_splitter.split_documents(docs)
# -----------------------
# Embedding Model
# -----------------------
embeddings = HuggingFaceEmbeddings(
model_name="all-MiniLM-L6-v2"
)
# -----------------------
# Vectorstore
# -----------------------
db = Chroma.from_documents(
documents=doc_chunks,
embedding=embeddings,
persist_directory="./chroma_db"
)
retriever = db.as_retriever(
search_type="similarity",
search_kwargs={"k": 1}
)
# -----------------------
# Prompt
# -----------------------
prompt_template = """
You are an AI assistant for Obiex.
Use ONLY the context below to answer the question.
Do NOT make assumptions. If the answer is not in the context, say 'I don't know'.
Context:
{context}
Question:
{question}
Answer:
"""
prompt = PromptTemplate.from_template(prompt_template)
# -----------------------
# LLM
# -----------------------
model_name = "microsoft/phi-4-mini-instruct"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(
model_name,
dtype=torch.float32,
)
stopping_criteria = StoppingCriteriaList([StopOnPunctuationOrMaxTokens()])
text_gen_pipe = pipeline(
"text-generation",
model=model,
tokenizer=tokenizer,
device=-1,
max_new_tokens=256,
temperature=0.1,
do_sample=False,
repetition_penalty=1.2,
return_full_text=False,
stopping_criteria=stopping_criteria
)
llm = HuggingFacePipeline(pipeline=text_gen_pipe, model_kwargs={"stop": ["Question:", "\n\n"]})
# -----------------------
# Output formatter
# -----------------------
def format_docs(documents):
return "\n\n".join(doc.page_content for doc in documents)
# -----------------------
# RAG chain
# -----------------------
rag_chain = (
{
"context": retriever | format_docs,
"question": RunnablePassthrough()
}
| prompt
| llm
| CleanStrOutputParser()
)
# -----------------------
# API endpoint
# -----------------------
@app.post("/query")
def query_rag(q: Query):
try:
response = rag_chain.invoke(q.question)
answer = response.split("Answer:")[-1].strip()
return {
"answer": answer,
}
except Exception as e:
print("ERROR:")
traceback.print_exc()
return {"error": str(e)}
@app.get("/health")
def health():
return {"status": "ok"}