ppo_test_space / app.py
ikenna1234's picture
Changes
0877002
Raw
History Blame Contribute Delete
4.12 kB
import gradio as gr
from typing import List, Union, Dict, Tuple
from transformers import pipeline
from os import getenv
from huggingface_hub import login
import pymupdf
from history import get_history, update_history
# Login to Hugging Face
login(getenv("Token"))
#name of model on huggingFace
model="ikenna1234/llama_3.2_1b_instruct_base_rlhf"
# Define generator pipeline
generator = pipeline("text-generation", model=model)
#Transform gradio history by breaking any tuple into 2 dicts
def transform_gradio_history(history: List[Union[Dict[str, str], Tuple[str, str]]]) -> List[Dict[str, str]]:
transformed_history = []
for entry in history:
if (isinstance(entry, list) or isinstance(entry, tuple)) and len(entry) == 2:
transformed_history.append({"role": "user", "content": entry[0]})
transformed_history.append({"role": "assistant", "content": entry[1]})
elif isinstance(entry, dict):
transformed_history.append(entry)
return transformed_history
def extract_text_from_pdf(pdf_path):
doc = pymupdf.open(pdf_path)
text = ""
for page in doc:
text += page.get_text()
return text
#Does the actual inference and streams (yield) the response
def chat(history:list[dict[str, str]],temperature,top_p,max_tokens,top_k):
for msg in generator(
history, #message list
max_new_tokens=max_tokens,
return_full_text=False,
temperature=temperature,
top_p=top_p,
top_k=top_k
#max_tokens=max_tokens
):
yield msg['generated_text']
def respond(
message,
history: list[dict[str, str]],
system_message, #system prompt
max_tokens,
temperature,
top_p,
file,
group_name #user Id
):
if not group_name:
#user must pass user Id to the group_name.
#This is used to identify the user
yield "User ID required"
else:
messages=history
#If no history, get history from database
if not len(messages):
messages=get_history(group_name)
#Break any tuples into 2 dicts
messages=transform_gradio_history(messages)
#Extract text from file
file_text=extract_text_from_pdf(file)
print("The file text: ", file_text)
#Add prompt to list of messages
messages.append({"role": "user", "content": message})
response = ""
#Create new list of all messages, starting with system prompt
mainMessage=[{"role": "system", "content": system_message}, *messages]
#calls the inference function and streams the response
for msg in chat(
mainMessage,
temperature=temperature,
top_p=top_p,
max_tokens=max_tokens,
top_k=12
):
token = msg
# This is a stream. Meaning response comes in bits of string.
# Add new response string bit to previous response
# strings to form the whole string
response += token
yield response
#update the history in database
if response:
messages.append({"role": "assistant", "content": response})
update_history(group_name,messages)
def initialize():
messages=[]
return messages
demo = gr.ChatInterface(
respond,
type="messages",
chatbot=gr.Chatbot(value=initialize(),type="messages"),
additional_inputs=[
gr.Textbox(value="You are an AI assistant that conducts interview", label="System message"),
gr.Slider(minimum=1, maximum=2048, value=512, step=1, label="Max new tokens"),
gr.Slider(minimum=0.1, maximum=4.0, value=0.7, step=0.1, label="Temperature"),
gr.Slider(
minimum=0.1,
maximum=1.0,
value=0.95,
step=0.05,
label="Top-p (nucleus sampling)",
),
gr.File(label="Upload File"),
gr.Textbox( label="User ID"),
],
)
if __name__ == "__main__":
demo.launch(share=True,ssr_mode=False)