CoolFace
Apppublic

Goodnight7/Medical_RLHF

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
utils.py113 linesDownload Raw Back to root
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        )