Goodnight7/Medical_RLHF
0
1# utils 2 3from langchain_chroma import Chroma4from langchain_nomic.embeddings import NomicEmbeddings5from langchain_core.documents import Document6from langchain.retrievers.document_compressors import CohereRerank7from langchain.retrievers import ContextualCompressionRetriever8from langchain.retrievers import EnsembleRetriever9from langchain_community.retrievers import BM25Retriever10from langchain_groq import ChatGroq11 12from dotenv import load_dotenv13from langchain_core.prompts import ChatPromptTemplate14from langchain_core.runnables import Runnable, RunnableMap15from langchain.schema import BaseRetriever16from qdrant_client import models17 18 19from langchain_huggingface.embeddings import HuggingFaceEmbeddings20 21load_dotenv()22#Retriever23def retriever(n_docs=5):24 vector_database_path = "chromadb3"25 26 #embeddings_model = NomicEmbeddings(model="nomic-embed-text-v1.5", inference_mode="local")27 embedding_model = HuggingFaceEmbeddings(model_name="sentence-transformers/all-mpnet-base-v2")28 29 30 vectorstore = Chroma(collection_name="chroma_db",31 persist_directory=vector_database_path,32 embedding_function=embedding_model)33 34 vs_retriever = vectorstore.as_retriever(k=n_docs)35 36 texts = vectorstore.get()['documents']37 metadatas = vectorstore.get()["metadatas"]38 39 documents = []40 for i in range(len(texts)):41 doc = Document(page_content=texts[i], metadata=metadatas[i])42 documents.append(doc)43 44 keyword_retriever = BM25Retriever.from_documents(documents)45 keyword_retriever.k = n_docs46 47 ensemble_retriever = EnsembleRetriever(retrievers=[vs_retriever,keyword_retriever],48 weights=[0.5, 0.5])49 50 compressor = CohereRerank(model="rerank-english-v3.0")51 retriever = ContextualCompressionRetriever(52 base_compressor=compressor, base_retriever=ensemble_retriever53 )54 55 return retriever56 57#Retriever prompt58rag_prompt = """You are a medical chatbot designed to answer health-related questions.59The questions you will receive will primarily focus on medical topics and patient care.60Here is the context to use to answer the question:61{context}62Think carefully about the above context.63Now, review the user question:64{input}65Provide an answer to this question using only the above context.66Answer:"""67 68# Post-processing69def format_docs(docs):70 return "\n\n".join(doc.page_content for doc in docs)71 72#RAG chain73def get_expression_chain(retriever: BaseRetriever, model_name="llama-3.1-70b-versatile", temp=0 ) -> Runnable:74 """Return a chain defined primarily in LangChain Expression Language"""75 def retrieve_context(input_text):76 # Use the retriever to fetch relevant documents77 docs = retriever.get_relevant_documents(input_text)78 return format_docs(docs)79 80 ingress = RunnableMap(81 {82 "input": lambda x: x["input"],83 "context": lambda x: retrieve_context(x["input"]),84 }85 )86 prompt = ChatPromptTemplate.from_messages(87 [88 (89 "system",90 rag_prompt91 )92 ]93 )94 llm = ChatGroq(model=model_name, temperature=temp)95 96 chain = ingress | prompt | llm97 return chain98 99#embedding_model = NomicEmbeddings(model="nomic-embed-text-v1.5", inference_mode="local")100embedding_model = HuggingFaceEmbeddings(model_name="sentence-transformers/all-mpnet-base-v2")101 102#Generate embeddings for a given text103def get_embeddings(text):104 return embedding_model.embed_query([text])[0] #, task_type='search_document'105 106 107# Create or connect to a Qdrant collection108def create_qdrant_collection(client, collection_name):109 if collection_name not in client.get_collections().collections:110 client.create_collection(111 collection_name=collection_name,112 vectors_config=models.VectorParams(size=768, distance=models.Distance.COSINE)113 )