CoolFace
Apppublic

VanguardAI/MultiModal_OpenSource_AI

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py434 linesDownload Raw Back to root
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()