srinit62/RAG_Implementation
0
1from langchain_huggingface import HuggingFaceEmbeddings2from langchain_chroma import Chroma3# Importing RunnableLambda for wrapping functions as runnable components.4from langchain_core.runnables import RunnableLambda5from langchain_core.prompts import PromptTemplate # To format prompts6from langchain_core.output_parsers import StrOutputParser # to transform the output of an LLM into a more usable format7from langchain.schema.runnable import RunnableParallel, RunnablePassthrough 8# importing HuggingFace model abstraction class from langchain9from langchain_huggingface import HuggingFaceEndpoint10 11import torch12import gradio13import re14 15# Initializing a Hugging Face endpoint for text generation using a specific model.16llm = HuggingFaceEndpoint(17 repo_id="HuggingFaceH4/zephyr-7b-beta", # Specifies the model to be used (Zephyr-7B Beta), available on Hugging Face.18 19 task="text-generation", # Defines the task type as text generation (commonly used for chatbots, summarization, etc.).20 21 max_new_tokens=512, # Limits the maximum number of new tokens the model can generate in a response.22 23 top_k=30, # Controls the diversity of the generated text by restricting the number of top probable tokens considered at each step.24 25 temperature=0.1, # A low temperature value makes the model’s output more deterministic and focused.26 27 repetition_penalty=1.03, # Slightly penalizes repeated tokens to reduce redundancy in responses.28)29 30# Initializing the Hugging Face embedding model.31embedding = HuggingFaceEmbeddings(32 model_name="mixedbread-ai/mxbai-embed-large-v1", # Using a pre-trained embedding model for text vectorization.33 34 model_kwargs={'device': "cuda" if torch.cuda.is_available() else "cpu"}, # Configuring the model to use GPU if available.35 36 encode_kwargs={'normalize_embeddings': False} # Disabling normalization to preserve the original embedding scale.37)38 39# Creating a Chroma vector database instance.40vectordb = Chroma(41 persist_directory= './chroma/', #'RAG_Implementation/chroma/', # Setting the directory where embeddings will be stored.42 embedding_function=embedding # Specifying the embedding function to convert text into vector representations.43)44 45# Function to retrieve relevant document chunks based on a given question.46def get_context_info(question):47 # Creating a retriever using the Chroma vector database with Maximum Marginal Relevance (MMR).48 # `search_type="mmr"` ensures diversity in retrieved results.49 # `fetch_k=5` first fetches the top 5 most relevant chunks.50 # `k=3` selects the 3 most diverse and relevant chunks from the 5 fetched.51 retriever = vectordb.as_retriever(search_type="mmr", search_kwargs={"k": 3, "fetch_k": 5})52 53 # Retrieving the document chunks using the retriever.54 docs = retriever.invoke(question)55 56 # Returning the retrieved document chunks.57 return docs58 59# Creating a retrieval pipeline using RunnableParallel.60retrieval = RunnableParallel(61 {62 # Retrieving relevant document chunks based on the input question.63 # RunnableLambda ensures that `get_context_info` runs dynamically when called.64 "context": RunnableLambda(lambda x: get_context_info(x["question"])),65 66 # Simply passing the question as is.67 "question": RunnableLambda(lambda x: x["question"])68 }69)70 71# Defining a prompt template for answering questions based on retrieved context.72template = """Use the following pieces of context to answer the question at the end.73If you don't know the answer, that is if the answer is not in the context, then just say that you don't know, don't try to make up an answer.74Always say "thanks for asking!" at the end of the answer.75 76{context}77Question: {question}78Helpful Answer:"""79 80# Creating a PromptTemplate object using the defined template.81# The template takes two input variables: 'context' (retrieved documents) and 'question' (user query).82QA_PROMPT = PromptTemplate(input_variables=["context", "question"], template=template)83 84# Constructing a Retrieval-Augmented Generation (RAG) chain.85 86rag_chain = (retrieval # Step 1: Retrieve relevant document chunks based on the user's question.87 | QA_PROMPT # Step 2: Format the retrieved context and question into a structured prompt.88 | llm # Step 3: Pass the prompt to the language model (LLM) for answer generation.89 | StrOutputParser() # Step 4: Convert the LLM's output into a plain string format.90 )91 92def clean_text(question):93 question=re.sub(r'\s+', ' ', question).strip()94 question=re.sub(r'\A\d+', '', question).strip()95 question = re.sub(r'#\S+', '', question)96 question = re.sub(r'@\S+', '', question)97 question.replace('\n', ' ')98 re.sub(r'\s+(?=(\n+$))', '', question)99 pattern = r'[^a-zA-Z0-9\s]'100 question=re.sub(pattern, '', question)101 return question102 103def get_rag_chain_response(user_question):104 user_question = clean_text(user_question)105 return rag_chain.invoke({"question": user_question})106 107#Below function is used only for testing. Otherwise it is not actually used108def get_context_string_testing(user_question):109 ret_arr = get_context_info(user_question)110 context_str = ''111 for i in range(len(ret_arr)):112 # Printing the content of each document chunk (page content).113 context_str += ret_arr[i].page_content + '\n ' 114 return context_str115 116# Gradio elements117 118# Input from user119in_max_length = 200 # YOUR CODE HERE120in_question = gradio.Textbox(type='text', label='Question', max_length=in_max_length)121 122# Output response123out_response = gradio.Textbox(type='text', label='Response to question')124 125# Gradio interface to generate UI link126iface = gradio.Interface( # YOUR CODE HERE127 fn=get_rag_chain_response, #get_context_string_testing,128 inputs=in_question,129 outputs=out_response,130 title="RAG implementation",131 description="via gradio",)132 133# YOUR CODE HERE to launch the interface134iface.launch()