| from fastapi import FastAPI, HTTPException |
| from pydantic import BaseModel |
|
|
| from langchain_community.document_loaders import PyPDFLoader |
| from langchain_text_splitters import RecursiveCharacterTextSplitter |
| |
| 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 |
|
|
| |
| |
| class CleanStrOutputParser(Runnable): |
| def invoke(self, input_text: str, config=None) -> str: |
| |
| text = input_text.strip() |
| if text.startswith('"') and text.endswith('"'): |
| text = text[1:-1] |
| if "." in text: |
| text = text[: text.rfind(".") + 1] |
| return text |
|
|
| |
| |
| |
| class Query(BaseModel): |
| question: str |
|
|
|
|
| |
| |
| |
|
|
| loader = PyPDFLoader("seplat.pdf") |
| docs = loader.load() |
|
|
| text_splitter = RecursiveCharacterTextSplitter( |
| chunk_size=1000, |
| chunk_overlap=200 |
| ) |
|
|
| doc_chunks = text_splitter.split_documents(docs) |
|
|
|
|
| |
| |
| |
|
|
| embeddings = HuggingFaceEmbeddings( |
| model_name="all-MiniLM-L6-v2" |
| ) |
|
|
|
|
| |
| |
| |
|
|
| 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_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) |
|
|
|
|
| |
| |
| |
|
|
| 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"]}) |
|
|
| |
| |
| |
|
|
| def format_docs(documents): |
| return "\n\n".join(doc.page_content for doc in documents) |
|
|
| |
| |
| |
|
|
| rag_chain = ( |
| { |
| "context": retriever | format_docs, |
| "question": RunnablePassthrough() |
| } |
| | prompt |
| | llm |
| | CleanStrOutputParser() |
| ) |
|
|
| |
| |
| |
|
|
| @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"} |