MuMoLLM / app.py
garima-mahato's picture
Update app.py
a96a463 verified
Raw
History Blame Contribute Delete
7.57 kB
import gradio as gr
import os
import time
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, StoppingCriteriaList, TextIteratorStreamer, pipeline
from threading import Thread
from peft import PeftModel
import whisperx
import gc
import pandas as pd
from transformers.utils import logging
logging.set_verbosity_info()
logger = logging.get_logger("transformers")
device = 'cuda' if torch.cuda.is_available() else 'cpu'
q_ctxt, query, ctxt_type = None, None, None
##################################################################
############################## Text ##############################
phi_base_model = AutoModelForCausalLM.from_pretrained(
'microsoft/phi-2',
low_cpu_mem_usage=True,
return_dict=True,
torch_dtype=torch.float32,
trust_remote_code=True
)
phi_new_model = "models/phi_adapter"
phi_model = PeftModel.from_pretrained(phi_base_model, phi_new_model)
phi_model = phi_model.merge_and_unload().to(device)
tokenizer = AutoTokenizer.from_pretrained('microsoft/phi-2')
tokenizer.pad_token = tokenizer.unk_token
###################################################################
############################## Audio ##############################
# 1. Transcribe with original whisper (batched)
audio_model_name = "small" #"large-v2"
audio_compute_type = "int8" if device == "cpu" else "float16" #"float16" # change to "int8" if low on GPU mem (may reduce accuracy)
audio_model = whisperx.load_model(audio_model_name, device, compute_type=audio_compute_type)
def get_audio_context(audio_file):
audio = whisperx.load_audio(audio_file)
batch_size = 4 if device == "cpu" else 16 #16 # reduce if low on GPU mem
result = audio_model.transcribe(audio, batch_size=batch_size)
# print(result["segments"]) # before alignment
context = " ".join(pd.DataFrame(result["segments"])["text"])
return context
###################################################################
#-----------------------------------------------------#
# Chatbot demo with multimodal input (text, markdown, LaTeX, code blocks, image, audio, & video). Plus shows support for streaming text.
def print_like_dislike(x: gr.LikeData):
print(x.index, x.value, x.liked)
def add_text(history, text):
history = history + [(text, None)]
return history, gr.Textbox(value="", interactive=False)
def add_file(history, file):
context = get_audio_context(file)
history = history + [((file.name,context), None)]
return history
def bot(history):
try:
global q_ctxt, query, ctxt_type
if len(history) > 0 and isinstance(history[-1][0], tuple):
ctxt_type = "audio"
q_ctxt = get_audio_context(history[-1][0][0])
response = "**Based on the given context:** \n" + q_ctxt + "\n. **What do you want to know ?**"
elif len(history) > 0 and isinstance(history[-1][0],str) and len(history[-1][0].strip()) > 0:
logger.info("Last History:" + history[-1][0])
if ctxt_type == "audio" and "[INST]" not in history[-1][0]:
query = history[-1][0]
input_text = q_ctxt + query
elif "[INST]" in history[-1][0]:
q_ctxt = history[-1][0].split("[INST]")[0]
query = history[-1][0].split("[INST]")[-1]
# input_text = history[-1][0].replace("[INST]"," ")
input_text = q_ctxt + query
ctxt_type = "text"
else:
input_text = q_ctxt + query
input_tokens = tokenizer.encode(input_text)
input_ids = torch.tensor(input_tokens, dtype=torch.int32).unsqueeze(0).to(device)
inputs_embeds = phi_model.get_input_embeddings()(input_ids)
out = phi_model.generate(inputs_embeds=inputs_embeds, min_new_tokens=10, max_new_tokens=50, bos_token_id=tokenizer.bos_token_id)
response = tokenizer.decode(out[0], skip_special_tokens=True)
else:
ctxt_type = None
q_ctxt = None
query = None
response = "Please ask your query or upload an audio/video to ask query"
logger.info(f"Context: {q_ctxt}")
logger.info(f"Query: {query}")
logger.info(f"Context Type: {ctxt_type}")
history[-1][1] = ""
for character in response:
history[-1][1] += character
time.sleep(0.05)
yield history
except Exception as e:
logger.error(e)
return e
def reset():
global q_ctxt, query, ctxt_type
q_ctxt, query, ctxt_type = None, None, None
return {
chatbot: None
}
## Style
title_markdown = ("""
<div align="center" background-color="#8bd4e2">
<h1 align="center"><a href="https://github.com/garima-mahato/MuMoLLM"><img src="https://raw.githubusercontent.com/garima-mahato/MuMoLLM/main/images/mumo_logo.png", alt="MuMo" border="0" style="margin: 0 auto; height: 200px;" /></a> </h1>
<h2 align="center"> MuMo: Multi-Modal Chatter </h2>
<h5 align="center"> If you like this project, please give a star ✨ on Github. </h2>
<h5 align="center"> Note: This version is not multilingual demo </h5>
<div align="center">
<div style="display:flex; gap: 0.25rem;" align="center">
<a href='https://github.com/garima-mahato/MuMoLLM'><img src='https://img.shields.io/badge/Github-Code-blue'></a>
<a href='https://github.com/garima-mahato/MuMoLLM/stargazers'><img src='https://img.shields.io/github/stars/garima-mahato/MuMoLLM.svg?style=social'></a>
<a href='https://www.youtube.com/watch?v=0DCZLFmnV0s'><img src='https://img.shields.io/badge/youtube-badge'></a>
</div>
</div>
</div>
\n**Notice**: This runs on CPU and may take time in responding.\n
**We recommend only one image or video per conversation session.** \n
If you want to start chatting with new images or videos, we recommend you to **CLEAR** the history to restart.\n
**For Text QnA: Please separate context and query by [INST]**
""")
css = """
pre {
white-space: pre-wrap; /* Since CSS 2.1 */
white-space: -moz-pre-wrap; /* Mozilla, since 1999 */
white-space: -pre-wrap; /* Opera 4-6 */
white-space: -o-pre-wrap; /* Opera 7 */
word-wrap: break-word; /* Internet Explorer 5.5+ */
}
"""
## Gradio App Set-up
with gr.Blocks(title="MuMo", css=css) as demo:
gr.Markdown(title_markdown)
chatbot = gr.Chatbot(
[],
elem_id="chatbot",
bubble_full_width=False,
#avatar_images=("https://raw.githubusercontent.com/garima-mahato/MuMoLLM/main/images/user_avatar.png", "https://raw.githubusercontent.com/garima-mahato/MuMoLLM/main/images/mumo_avatar.png"),
)
with gr.Row():
txt = gr.Textbox(
scale=4,
show_label=False,
placeholder="Enter text and press enter, or upload an image",
container=False,
)
btn = gr.UploadButton("📁", file_types=["image", "video", "audio"])
txt_msg = txt.submit(add_text, [chatbot, txt], [chatbot, txt], queue=False).then(
bot, chatbot, chatbot, api_name="bot_response"
)
txt_msg.then(lambda: gr.Textbox(interactive=True), None, [txt], queue=False)
file_msg = btn.upload(add_file, [chatbot, btn], [chatbot], queue=False).then(
bot, chatbot, chatbot
)
with gr.Row():
clear = gr.Button("Clear")
chatbot.like(print_like_dislike, None, None)
clear.click(reset, None, chatbot, queue=False)
demo.queue()
demo.launch(debug=True)