CoolFace
Apppublic

G1GenAI/LagRAG_demo

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py133 linesDownload Raw Back to root
1import streamlit as st2import torch3from transformers import pipeline, AutoTokenizer, AutoModelForCausalLM4from transformers import StoppingCriteriaList, StoppingCriteria5from sentence_transformers import SentenceTransformer6from pinecone import Pinecone7import warnings8 9 10warnings.filterwarnings("ignore", category=UserWarning)11 12# model_name = "AI-Sweden-Models/gpt-sw3-126m-instruct"13model_name = "AI-Sweden-Models/gpt-sw3-126m-instruct"14 15 16device = "cuda:0" if torch.cuda.is_available() else "cpu"17 18# Initialize Tokenizer & Model19tokenizer = AutoTokenizer.from_pretrained(model_name)20 21 22def read_file(file_path: str) -> str:23    """Read the contents of a file."""24    with open(file_path, "r") as file:25        return file.read()26 27 28model = AutoModelForCausalLM.from_pretrained(model_name)29model.eval()30model.to(device)31 32document_encoder_model = SentenceTransformer("KBLab/sentence-bert-swedish-cased")33 34 35# Note: 'index1' has been pre-created in the pinecone console36# read the pinecone api key from a file37pinecone_api_key = st.secrets["pinecone_api_key"]38pc = Pinecone(api_key=pinecone_api_key)39index = pc.Index("index1")40 41 42def query_pincecone_namespace(43    vector_databse_index: Pinecone, q_embedding: str, namespace: str44) -> str:45    result = vector_databse_index.query(46        namespace=namespace,47        vector=q_embedding.tolist(),48        top_k=1,49        include_values=True,50        include_metadata=True,51    )52    results = []53    for match in result.matches:54        results.append(match.metadata["paragraph"])55    return results[0]56 57 58def generate_prompt(llmprompt: str) -> str:59    """Generates a prompt for the GPT-3 model"""60    start_token = "<|endoftext|><s>"61    end_token = "<s>"62    return f"{start_token}\nUser:\n{llmprompt}\n{end_token}\nBot:\n".strip()63 64 65def encode_query(query: str) -> torch.Tensor:66    """Encode the query using the model's tokenizer"""67    return document_encoder_model.encode(query)68 69 70class StopOnTokenCriteria(StoppingCriteria):71    def __init__(self, stop_token_id):72        self.stop_token_id = stop_token_id73 74    def __call__(self, input_ids, scores, **kwargs):75        return input_ids[0, -1] == self.stop_token_id76 77 78stop_on_token_criteria = StopOnTokenCriteria(stop_token_id=tokenizer.bos_token_id)79 80st.title("Paralegal Assistant")81st.subheader("RAG: föräldrabalken")82 83# Initialize chat history84if "messages" not in st.session_state:85    st.session_state.messages = []86 87# Display chat messages from history on app rerun88for message in st.session_state.messages:89    with st.chat_message(message["role"]):90        st.markdown(message["content"])91 92# React to user input93if prompt := st.chat_input("Skriv din fråga..."):94    # Display user message in chat message container95    st.chat_message("user").markdown(prompt)96    # Add user message to chat history97    st.session_state.messages.append({"role": "user", "content": prompt})98 99    query = query_pincecone_namespace(100        vector_databse_index=index,101        q_embedding=encode_query(query=prompt),102        namespace="ns-parent-balk",103    )104    llmprompt = (105        "Följande stycke är en del av lagen: "106        + query107        +"Referera till lagen och besvara följande fråga på ett sakligt, kortfattat och formellt vis: "108        + prompt109    )110    llmprompt = generate_prompt(llmprompt=llmprompt)111 112    # # Convert prompt to tokens113    input_ids = tokenizer(llmprompt, return_tensors="pt")["input_ids"].to(device)114 115    # Genqerate tokens based om prompt116    generated_token_ids = model.generate(117        inputs=input_ids,118        max_new_tokens=128,119        do_sample=True,120        temperature=0.8,121        top_p=1,122        stopping_criteria=StoppingCriteriaList([stop_on_token_criteria]),123    )[0]124 125    # Decode the generated tokens126    generated_text = tokenizer.decode(generated_token_ids[len(input_ids[0]) : -1])127 128    response = f"{generated_text}"129    # Display assistant response in chat message container130    with st.chat_message("assistant"):131        st.markdown(f"```{query}```\n" + response)132    # Add assistant response to chat history133    st.session_state.messages.append({"role": "assistant", "content": response})