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"}