Samizie/Chat_RAG
0
1import asyncio2import requests3import pandas as pd4import re5import numpy as np6import faiss7from langchain_community.document_loaders import AsyncChromiumLoader8from langchain_community.document_transformers import Html2TextTransformer9from langchain.text_splitter import RecursiveCharacterTextSplitter10from langchain_ollama import OllamaLLM11#from langchain_ollama import OllamaEmbeddings12from langchain_groq import ChatGroq13from itertools import chain14from sentence_transformers import SentenceTransformer15from langchain_community.vectorstores import FAISS16from langchain.text_splitter import RecursiveCharacterTextSplitter17 18 19async def process_urls(urls):20 # Load multiple URLs asynchronously21 loader = AsyncChromiumLoader(urls)22 docs = await loader.aload()23 24 # Transform HTML to text25 text_transformer = Html2TextTransformer()26 transformed_docs = text_transformer.transform_documents(docs)27 28 # Split the text into chunks and retain metadata29 text_splitter = RecursiveCharacterTextSplitter(chunk_size=5000, chunk_overlap=500)30 split_docs_nested = [text_splitter.split_documents([doc]) for doc in transformed_docs]31 #split_docs = text_splitter.split_documents(transformed_docs)32 split_docs = list(chain.from_iterable(split_docs_nested))33 # Attach the source URL to each split document34 for doc in split_docs:35 doc.metadata["source_url"] = doc.metadata.get("source", "Unknown") # Ensure URL metadata exists36 37 return split_docs38 39def clean_text(text):40 """Remove unnecessary whitespace, line breaks, and special characters."""41 text = re.sub(r'\s+', ' ', text).strip() # Remove excessive whitespace42 text = re.sub(r'\[.*?\]|\(.*?\)', '', text) # Remove bracketed text (e.g., [advert])43 return text44 45 46def embed_text(text_list):47 embeddings = SentenceTransformer("nomic-ai/nomic-embed-text-v1", trust_remote_code=True)48 #return embeddings.encode(text_list)49 if embeddings is None or len(embeddings) == 0:50 raise ValueError("Embedding function returned an empty result.")51 return embeddings.encode(text_list)52 53 54def store_embeddings(docs):55 """Convert text into embeddings and store them in FAISS."""56 #all_text = [clean_text(doc.page_content) for doc in docs if doc.page_content]57 all_text = [clean_text(doc.page_content) for doc in docs if hasattr(doc, "page_content")]58 text_sources = [doc.metadata["source_url"] for doc in docs]59 60 embeddings = embed_text(all_text)61 if embeddings is None or embeddings.size == 0:62 raise ValueError("Embedding function returned None or empty list.")63 64 embeddings = np.array(embeddings, dtype=np.float32)65 # Normalize embeddings for better FAISS similarity search66 faiss.normalize_L2(embeddings)67 d = embeddings.shape[1]68 index = faiss.IndexFlatIP(d) # Inner Product (cosine similarity)69 index.add(embeddings)70 71 return index, all_text, text_sources72 73def search_faiss(index, query_embedding, text_data, text_sources, top_k=5, min_score=0.5):74 #query_embedding = np.array([query_embedding], dtype=np.float32)75 query_embedding = query_embedding.reshape(1, -1) 76 faiss.normalize_L2(query_embedding) # Normalize query embedding for similarity77 78 distances, indices = index.search(query_embedding, top_k)79 80 results = []81 if indices.size > 0:82 for i in range(len(indices[0])):83 if distances[0][i] >= min_score: # Ignore irrelevant results84 idx = indices[0][i]85 if idx < len(text_data):86 results.append({"source": text_sources[idx], "content": text_data[idx]})87 88 return results89 90def query_llm(index, text_data, text_sources, query):91 groq_api="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"92 chat = ChatGroq(model="llama-3.2-1b-preview", groq_api_key=groq_api, temperature=0)93 94 # Embed the query95 query_embedding = embed_text([query])[0]96 97 # Search FAISS for relevant documents98 relevant_docs = search_faiss(index, query_embedding, text_data, text_sources, top_k=3)99 print(type(relevant_docs))100 print(relevant_docs)101 102 # If no relevant docs, return a default message103 if not relevant_docs:104 return "No relevant information found."105 106 # Query LLM with retrieved content107 responses = []108 for doc in relevant_docs:109 if isinstance(doc, dict) and "source" in doc and "content" in doc:110 source_url = doc["source"]111 content = doc["content"][:10000]112 else:113 print(f"Unexpected doc format: {doc}") # Debugging print114 continue115 116 prompt = f"""117 Based on the following content, answer the question: "{query}"118 119 Content (from {source_url}):120 {content}121 122 "123 """124 response = chat.invoke(prompt)125 #print(type(response))126 responses.append({"source": source_url, "response": response})127 128 return responses129 130 131 132 133#urls = ["https://edition.cnn.com/", "https://www.bbc.com/", "https://www.vanguardngr.com/"]134 135# query = "Where is Nigeria located"136 137# async def main():138# urls = ["https://en.wikipedia.org/wiki/Nigeria","https://en.wikipedia.org/wiki/Ghana"] # Replace with actual URLs139# split_docs = await process_urls(urls)140# print(split_docs)141 142 143#split_docs = process_urls(urls)144#print(split_docs)145#print(split_docs)146# index, text_data, text_sources = store_embeddings(split_docs)147# query_embedding = np.array([embed_text([query])[0]])#, dtype=np.float32)148# query_embedding = query_embedding.reshape(1, -1)149#print(split_docs[0].page_content)150#print(index)151#print(text_data)152#print(text_sources)153# response = query_llm(index, text_data, text_sources, query)154# print(response)155 156 157#print(query_embedding.shape)158#relevant_docs = search_faiss(index, query_embedding, text_data, text_sources, top_k=3)159#print(relevant_docs)160 161 162 163"""query = "Where is Nigeria located"164 165async def main():166 urls = ["https://en.wikipedia.org/wiki/Nigeria", "https://en.wikipedia.org/wiki/Ghana"]167 split_docs = await process_urls(urls)168 169 # Ensure split_docs is available before using it170 index, text_data, text_sources = store_embeddings(split_docs)171 172 query_embedding = np.array([embed_text([query])[0]])173 query_embedding = query_embedding.reshape(1, -1)174 175 response = query_llm(index, text_data, text_sources, query)176 print(response)177 178# Run the async function179asyncio.run(main())"""