Team Ai
Apppublic

amanm10000/MLSC-Coherence-25-FAQ-Chatbot-API

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
main.py160 linesDownload Raw Back to root
1import os2import getpass3from groq import Groq4from langchain.chat_models import init_chat_model5from langchain_core.messages import HumanMessage, SystemMessage6from langchain_core.vectorstores import InMemoryVectorStore7from langchain_core.documents import Document8from langchain_text_splitters import RecursiveCharacterTextSplitter9from langchain_community.document_loaders import UnstructuredMarkdownLoader10from langchain_community.embeddings import HuggingFaceInferenceAPIEmbeddings11from langchain import hub12from langgraph.graph import START, StateGraph13from pydantic.main import BaseModel14from typing_extensions import List, TypedDict15 16from langchain_cohere import CohereEmbeddings17 18import re19# from dotenv import load_dotenv20from fastapi import FastAPI21from fastapi.middleware.cors import CORSMiddleware22from fastapi.responses import JSONResponse23 24'''25if not os.environ.get("GROQ_API_KEY"):26    os.environ["GROQ_API_KEY"] = getpass.getpass("Enter API key for Groq: ")27'''28 29# load_dotenv()30 31# print(f"GROQ_API_KEY: {os.getenv('GROQ_API_KEY')}")32# print(f"HUGGING_FACE_API_KEY: {os.getenv('HUGGING_FACE_API_KEY')}")33 34llm = init_chat_model("qwen-qwq-32b", model_provider="groq", api_key=os.environ["GROQ_API_KEY"])35'''36embeddings = HuggingFaceInferenceAPIEmbeddings(37    api_key = os.getenv('HUGGING_FACE_API_KEY'),38    model_name="sentence-transformers/all-MiniLM-L6-v2"39)40 41embeddings = HuggingFaceInferenceAPIEmbeddings(42    api_key=os.getenv('HUGGING_FACE_API_KEY'), model_name="sentence-transformers/all-MiniLM-L6-v2"43)'''44 45embeddings = CohereEmbeddings(46    cohere_api_key=os.environ['COHERE'],47    model="embed-english-v3.0",  # Added this line48    user_agent="langchain-cohere-embeddings"49)50 51vector_store = InMemoryVectorStore(embedding=embeddings)52 53# Data - 1 and Data - 254data_1 = open(r'data_1.txt', 'r').read()55data_2 = open(r'data_2.txt', 'r').read()56data_3 = open(r'data_3.txt', 'r').read()57data_4 = open(r'data_4.txt', 'r').read()58 59comb = open(r'comb.txt', 'r').read()60 61md_loader = UnstructuredMarkdownLoader('comb.md')62 63text_splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=100)64# all_splits = text_splitter.split_text(data_1 + "\n\n" + data_2 + "\n\n" + data_3 + "\n\n" + data_4)65# all_splits = text_splitter.split_text(comb)66all_splits = text_splitter.split_documents(md_loader.load())67 68# docs = [Document(page_content=text) for text in all_splits]69docs = [Document(page_content=text.page_content, metadata=text.metadata) for text in all_splits]70_ = vector_store.add_documents(documents=docs)71 72 73prompt = hub.pull("rlm/rag-prompt")74 75# Replace with custom prompt76system_message = """You are a helpful and professional FAQ chatbot for the MLSC Coherence 25 Hackathon. Your role is to:771. Provide accurate and concise answers based on the provided context782. Be friendly but professional in tone793. If you don't know the answer, simply say "I don't have information about that"804. Keep responses brief and to the point815. Focus on providing factual information from the context826. Never mention "the provided context" or similar phrases in your responses837. Never explain why you don't know something - just state that you don't know848. Be direct and avoid unnecessary explanations"""85 86human_message_template = """Context: {context}87 88Question: {question}89 90Please provide a clear and concise answer based on the context above."""91 92class State(TypedDict):93    question: str94    context: List[Document]95    answer: str96 97def retrieve(state: State):98    retrieved_docs = vector_store.similarity_search(state["question"])99    return {"context": retrieved_docs}100 101def generate(state: State):102    docs_content = "\n\n".join(doc.page_content for doc in state["context"])103    messages = [104        SystemMessage(content=system_message),105        HumanMessage(content=human_message_template.format(106            context=docs_content,107            question=state["question"]108        ))109    ]110    print(messages)111    response = llm.invoke(messages)112    return {"answer": response.content}113 114graph_builder = StateGraph(State).add_sequence([retrieve, generate])115graph_builder.add_edge(START, "retrieve")116graph = graph_builder.compile()117'''118response = graph.invoke({"question": "Who should i contact for help ?"})119print(response["answer"])120'''121 122app = FastAPI()123 124origins = ["*"]125 126app.add_middleware(127    CORSMiddleware,128    allow_origins=origins,129    allow_credentials=True,130    allow_methods=["GET", "POST", "PUT", "DELETE"],131    allow_headers=["*"],132)133 134@app.get("/ping")135async def ping():136    return "Pong!"137 138class Query(BaseModel):139    question: str140 141@app.get("/chat")142async def chat(request: Query):143    response = graph.invoke({"question": request.question})144    response = response["answer"]145    response = re.sub(r'<think>.*?</think>', '', response, flags=re.DOTALL)146    # response = response[4:]147    return {"response": response}148 149@app.post("/chat")150async def chat(request: Query):151    response = graph.invoke({"question": request.question})152    response = response["answer"]153    response = re.sub(r'<think>.*?</think>', '', response, flags=re.DOTALL)154    # response = response[4:]155    return {"response": response}156 157@app.get("/")158async def root():159    return {"message": "Hello World"}160