G1GenAI/LagRAG_demo
0
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})