VanguardAI/MultiModal_OpenSource_AI
0
1import gradio as gr2import torch3import os4import numpy as np5from groq import Groq6import spaces7from transformers import AutoModel, AutoTokenizer8from diffusers import StableDiffusion3Pipeline9from parler_tts import ParlerTTSForConditionalGeneration10import soundfile as sf11from langchain_groq import ChatGroq12from PIL import Image13from tavily import TavilyClient14from langchain.schema import AIMessage15from langchain_community.embeddings import HuggingFaceEmbeddings16from langchain_community.vectorstores import FAISS17from langchain_community.document_loaders import TextLoader18from langchain.text_splitter import CharacterTextSplitter19from langchain.chains import RetrievalQA20from torchvision import transforms21import json22import pandas23 24# Initialize models and clients25MODEL = 'llama-3.1-70b-versatile'26client = Groq(api_key=os.environ.get("GROQ_API_KEY"))27 28vqa_model = AutoModel.from_pretrained('openbmb/MiniCPM-V-2', trust_remote_code=True,29 device_map="auto", torch_dtype=torch.bfloat16)30tokenizer = AutoTokenizer.from_pretrained('openbmb/MiniCPM-V-2', trust_remote_code=True)31 32tts_model = ParlerTTSForConditionalGeneration.from_pretrained("parler-tts/parler-tts-large-v1")33tts_tokenizer = AutoTokenizer.from_pretrained("parler-tts/parler-tts-large-v1")34 35# Updated Image generation model36pipe = StableDiffusion3Pipeline.from_pretrained("stabilityai/stable-diffusion-3-medium-diffusers", torch_dtype=torch.float16)37pipe = pipe.to("cuda")38 39# Tavily Client for web search40tavily_client = TavilyClient(api_key=os.environ.get("TAVILY_API"))41 42# Function to play voice output43def play_voice_output(response):44 print("Executing play_voice_output function")45 description = "Jon's voice is monotone yet slightly fast in delivery, with a very close recording that almost has no background noise."46 input_ids = tts_tokenizer(description, return_tensors="pt").input_ids.to('cuda')47 prompt_input_ids = tts_tokenizer(response, return_tensors="pt").input_ids.to('cuda')48 generation = tts_model.generate(input_ids=input_ids, prompt_input_ids=prompt_input_ids)49 audio_arr = generation.cpu().numpy().squeeze()50 sf.write("output.wav", audio_arr, tts_model.config.sampling_rate)51 return "output.wav"52 53# Function to classify user input using LLM54def classify_function(user_prompt):55 prompt = f"""56 You are a function classifier AI assistant. You are given a user input and you need to classify it into one of the following functions:57 58 - `image_generation`: If the user wants to generate an image.59 - `image_vqa`: If the user wants to ask questions about an image.60 - `document_qa`: If the user wants to ask questions about a document.61 - `text_to_text`: If the user wants a text-based response.62 63 Respond with a JSON object containing only the chosen function. For example:64 65 ```json66 {{"function": "image_generation"}}67 ```68 69 User input: {user_prompt}70 """71 72 chat_completion = client.chat.completions.create(73 messages=[74 {75 "role": "user",76 "content": prompt,77 }78 ],79 model="llama3-8b-8192",80 )81 82 try:83 response = json.loads(chat_completion.choices[0].message.content)84 function = response.get("function")85 return function86 except json.JSONDecodeError:87 print(f"Error decoding JSON: {chat_completion.choices[0].message.content}")88 return "text_to_text" # Default to text-to-text if JSON parsing fails89 90# Document Question Answering Tool91class DocumentQuestionAnswering:92 def __init__(self, document):93 self.document = document94 self.qa_chain = self._setup_qa_chain()95 96 def _setup_qa_chain(self):97 print("Setting up DocumentQuestionAnswering tool")98 loader = TextLoader(self.document)99 documents = loader.load()100 text_splitter = CharacterTextSplitter(chunk_size=1000, chunk_overlap=0)101 texts = text_splitter.split_documents(documents)102 embeddings = HuggingFaceEmbeddings()103 db = FAISS.from_documents(texts, embeddings)104 retriever = db.as_retriever()105 qa_chain = RetrievalQA.from_chain_type(106 llm=ChatGroq(model=MODEL, api_key=os.environ.get("GROQ_API_KEY")),107 chain_type="stuff",108 retriever=retriever,109 )110 return qa_chain111 112 def run(self, query: str) -> str:113 print("Executing DocumentQuestionAnswering tool")114 response = self.qa_chain.run(query)115 return str(response)116 117# Function to handle different input types and choose the right pipeline118def handle_input(user_prompt, image=None, audio=None, websearch=False, document=None):119 print(f"Handling input: {user_prompt}")120 121 # Initialize the LLM122 llm = ChatGroq(model=MODEL, api_key=os.environ.get("GROQ_API_KEY"))123 124 # Handle voice-only mode125 if audio:126 print("Processing audio input")127 transcription = client.audio.transcriptions.create(128 file=(audio.name, audio.read()),129 model="whisper-large-v3"130 )131 user_prompt = transcription.text132 response = llm.invoke(query=user_prompt)133 audio_output = play_voice_output(response)134 return "Response generated.", audio_output135 136 # Handle websearch mode137 if websearch:138 print("Executing Web Search")139 answer = tavily_client.qna_search(query=user_prompt)140 return answer, None141 142 # Handle cases with only image or document input143 if user_prompt is None or user_prompt.strip() == "":144 if image:145 user_prompt = "Describe this image"146 elif document:147 user_prompt = "Summarize this document"148 149 # Classify user input using LLM150 function = classify_function(user_prompt)151 152 # Handle different functions153 if function == "image_generation":154 print("Executing Image Generation")155 image = pipe(156 user_prompt,157 negative_prompt="",158 num_inference_steps=15,159 guidance_scale=7.0,160 ).images[0]161 image.save("output.jpg")162 return "output.jpg", None163 164 elif function == "image_vqa":165 print("Executing Image Description")166 if image:167 print("1")168 image = Image.open(image).convert('RGB')169 print("2")170 171 # Add preprocessing steps here (see examples above)172 preprocess = transforms.Compose([173 transforms.Resize((512, 512)), # Example size, replace with the correct one174 transforms.ToTensor(),175 ])176 image = preprocess(image)177 image = image.unsqueeze(0) # Add batch dimension178 image = image.to(torch.float32) # Ensure correct data type179 180 print("3")181 messages = [{"role": "user", "content": user_prompt}]182 print("4")183 response,ctxt = vqa_model.chat(image=image, msgs=messages, tokenizer=tokenizer, context=None, temperature=0.5)184 print("5")185 return response, None186 else:187 return "Please upload an imagee.", None188 189 elif function == "document_qa":190 print("Executing Document Summarization")191 if document:192 document_qa = DocumentQuestionAnswering(document)193 response = document_qa.run(user_prompt)194 return response, None195 else:196 return "Please upload a documentt.", None197 198 else: # function == "text_to_text"199 print("Executing Text-to-Text")200 response = llm.invoke(query=user_prompt)201 return response, None202 203# Main interface function204@spaces.GPU(duration=120)205def main_interface(user_prompt, image=None, audio=None, voice_only=False, websearch=False, document=None):206 print("Starting main_interface function")207 vqa_model.to(device='cuda', dtype=torch.bfloat16)208 tts_model.to("cuda")209 pipe.to("cuda")210 211 print(f"user_prompt: {user_prompt}, image: {image}, audio: {audio}, voice_only: {voice_only}, websearch: {websearch}, document: {document}")212 213 try:214 response = handle_input(user_prompt, image=image, audio=audio, websearch=websearch, document=document)215 print("handle_input function executed successfully")216 except Exception as e:217 print(f"Error in handle_input: {e}")218 response = "Error occurred during processing."219 220 return response221 222def create_ui():223 with gr.Blocks(css="""224 /* Overall Styling */225 body {226 font-family: 'Poppins', sans-serif;227 background: linear-gradient(135deg, #f5f7fa 0%, #c3cfe2 100%);228 margin: 0;229 padding: 0;230 color: #333;231 }232 233 /* Title Styling */234 .gradio-container h1 {235 text-align: center;236 padding: 20px 0;237 background: linear-gradient(45deg, #007bff, #00c6ff);238 color: white;239 font-size: 2.5em;240 font-weight: bold;241 letter-spacing: 1px;242 text-transform: uppercase;243 margin: 0;244 box-shadow: 0px 4px 8px rgba(0, 0, 0, 0.2);245 }246 247 /* Input Area Styling */248 .gradio-container .gr-row {249 display: flex;250 justify-content: space-around;251 align-items: center;252 padding: 20px;253 background-color: white;254 border-radius: 10px;255 box-shadow: 0px 6px 12px rgba(0, 0, 0, 0.1);256 margin-bottom: 20px;257 }258 259 .gradio-container .gr-column {260 flex: 1;261 margin: 0 10px;262 }263 264 /* Textbox Styling */265 .gradio-container textarea {266 width: calc(100% - 20px);267 padding: 15px;268 border: 2px solid #007bff;269 border-radius: 8px;270 font-size: 1.1em;271 transition: border-color 0.3s, box-shadow 0.3s;272 }273 274 .gradio-container textarea:focus {275 border-color: #00c6ff;276 box-shadow: 0px 0px 8px rgba(0, 198, 255, 0.5);277 outline: none;278 }279 280 /* Button Styling */281 .gradio-container button {282 background: linear-gradient(45deg, #007bff, #00c6ff);283 color: white;284 padding: 15px 25px;285 border: none;286 border-radius: 8px;287 cursor: pointer;288 font-size: 1.2em;289 font-weight: bold;290 transition: background 0.3s, transform 0.3s;291 box-shadow: 0px 4px 8px rgba(0, 0, 0, 0.1);292 }293 294 .gradio-container button:hover {295 background: linear-gradient(45deg, #0056b3, #009bff);296 transform: translateY(-3px);297 }298 299 .gradio-container button:active {300 transform: translateY(0);301 }302 303 /* Output Area Styling */304 .gradio-container .output-area {305 padding: 20px;306 text-align: center;307 background-color: #f7f9fc;308 border-radius: 10px;309 box-shadow: 0px 6px 12px rgba(0, 0, 0, 0.1);310 margin-top: 20px;311 }312 313 /* Image Styling */314 .gradio-container img {315 max-width: 100%;316 height: auto;317 border-radius: 10px;318 box-shadow: 0px 4px 8px rgba(0, 0, 0, 0.1);319 transition: transform 0.3s, box-shadow 0.3s;320 }321 322 .gradio-container img:hover {323 transform: scale(1.05);324 box-shadow: 0px 6px 12px rgba(0, 0, 0, 0.2);325 }326 327 /* Checkbox Styling */328 .gradio-container input[type="checkbox"] {329 width: 20px;330 height: 20px;331 cursor: pointer;332 accent-color: #007bff;333 transition: transform 0.3s;334 }335 336 .gradio-container input[type="checkbox"]:checked {337 transform: scale(1.2);338 }339 340 /* Audio and Document Upload Styling */341 .gradio-container .gr-file-upload input[type="file"] {342 width: 100%;343 padding: 10px;344 border: 2px solid #007bff;345 border-radius: 8px;346 cursor: pointer;347 background-color: white;348 transition: border-color 0.3s, background-color 0.3s;349 }350 351 .gradio-container .gr-file-upload input[type="file"]:hover {352 border-color: #00c6ff;353 background-color: #f0f8ff;354 }355 356 /* Advanced Tooltip Styling */357 .gradio-container .gr-tooltip {358 position: relative;359 display: inline-block;360 cursor: pointer;361 }362 363 .gradio-container .gr-tooltip .tooltiptext {364 visibility: hidden;365 width: 200px;366 background-color: black;367 color: #fff;368 text-align: center;369 border-radius: 6px;370 padding: 5px;371 position: absolute;372 z-index: 1;373 bottom: 125%;374 left: 50%;375 margin-left: -100px;376 opacity: 0;377 transition: opacity 0.3s;378 }379 380 .gradio-container .gr-tooltip:hover .tooltiptext {381 visibility: visible;382 opacity: 1;383 }384 385 /* Footer Styling */386 .gradio-container footer {387 text-align: center;388 padding: 10px;389 background: #007bff;390 color: white;391 font-size: 0.9em;392 border-radius: 0 0 10px 10px;393 box-shadow: 0px -2px 8px rgba(0, 0, 0, 0.1);394 }395 396 """) as demo:397 gr.Markdown("# AI Assistant")398 with gr.Row():399 with gr.Column(scale=2):400 user_prompt = gr.Textbox(placeholder="Type your message here...", lines=1)401 with gr.Column(scale=1):402 image_input = gr.Image(type="filepath", label="Upload an image", elem_id="image-icon")403 audio_input = gr.Audio(type="filepath", label="Upload audio", elem_id="mic-icon")404 document_input = gr.File(type="filepath", label="Upload a document", elem_id="document-icon")405 voice_only_mode = gr.Checkbox(label="Enable Voice Only Mode", elem_id="voice-only-mode")406 websearch_mode = gr.Checkbox(label="Enable Web Search", elem_id="websearch-mode")407 with gr.Column(scale=1):408 submit = gr.Button("Submit")409 410 output_label = gr.Label(label="Output")411 audio_output = gr.Audio(label="Audio Output", visible=False)412 413 submit.click(414 fn=main_interface,415 inputs=[user_prompt, image_input, audio_input, voice_only_mode, websearch_mode, document_input],416 outputs=[output_label, audio_output]417 )418 419 voice_only_mode.change(420 lambda x: gr.update(visible=not x),421 inputs=voice_only_mode,422 outputs=[user_prompt, image_input, websearch_mode, document_input, submit]423 )424 voice_only_mode.change(425 lambda x: gr.update(visible=x),426 inputs=voice_only_mode,427 outputs=[audio_input]428 )429 430 return demo431 432# Launch the UI433demo = create_ui()434demo.launch()