import gradio as gr import tensorflow as tf from transformers import AutoImageProcessor, AutoModelForImageClassification import numpy as np import pymongo import random import base64 # Import base64 module import io # Import io module from PIL import Image # Import PIL print(gr.__version__) # --- Loading models and DB connection (Keep as is) --- processor = AutoImageProcessor.from_pretrained("SowmyaKannan/fine_tuned_vit_model") model = AutoModelForImageClassification.from_pretrained("SowmyaKannan/Vit-FER") labels = ['angry', 'disgust', 'fear', 'happy', 'sad', 'neutral', 'surprise'] mongo_db_uri = "mongodb+srv://sowmyakannan10:O2BkajONKtfZsNqm@cluster0.zehazwd.mongodb.net/?retryWrites=true&w=majority&appName=Cluster0" client = None db = None quotes_collection = None author_list = [] try: mongodb_uri = mongo_db_uri if not mongodb_uri: print("MONGODB_URI not found!!") else: print("Connecting to MongoDB...") client = pymongo.MongoClient(mongodb_uri, serverSelectionTimeoutMS=5000) client.admin.command('ping') print("Successfully connected to MongoDB!") db = client['AhumDB'] quotes_collection = db['AhumQuotes'] author_list = sorted(filter(None, quotes_collection.distinct("author"))) print(f"Fetched {len(author_list)} authors.") except Exception as e: print(f"Error during MongoDB setup: {e}") client = None # --- End loading models and DB connection --- # --- Quote fetching function (Keep as is) --- def get_quotes(index=None, authors=None): if client is None or quotes_collection is None: print("MongoDB connection not available. Cannot fetch quote.") return "Quote data unavailable (DB connection issue).", "Unknown" pipeline = [] emotion_label = None emotion_match = {} if index is not None and 0 <= index < len(labels): emotion_label = labels[index] emotion_match = {"emotions": {"$regex": f"^{emotion_label}$", "$options": "i"}} author_match = {} if authors: author_match = {"$or": [{"author": {"$regex": f"^{a}$", "$options": "i"}} for a in authors]} combined_match = {} if emotion_match and author_match: combined_match = {"$and": [emotion_match, author_match]} elif emotion_match: combined_match = emotion_match elif author_match: combined_match = author_match if combined_match: pipeline.append({"$match": combined_match}) pipeline.append({"$sample": {"size": 1}}) try: result = list(quotes_collection.aggregate(pipeline)) if result: return result[0].get('quote', "Quote not found."), result[0].get('author', "Unknown") else: # Fallback logic if emotion_label: fallback_pipeline = [{"$match": {"emotions": {"$regex": f"^{emotion_label}$", "$options": "i"}}}, {"$sample": {"size": 1}}] fallback_result = list(quotes_collection.aggregate(fallback_pipeline)) if fallback_result: return fallback_result[0].get('quote', "Quote not found."), fallback_result[0].get('author', "Unknown") return f"Finding strength in '{emotion_label}' comes from within.", "Unknown" else: return "Could not find a suitable quote at this time.", "Unknown" except Exception as e: print(f"Error in get_quotes: {e}") return "Quote data unavailable (DB query error).", "Unknown" # --- End quote fetching --- # --- CORE LOGIC: Process PIL image and get prediction/quote --- def process_pil_image_and_get_quote(img_pil, selected_authors): """Takes a PIL image and author list, returns formatted prediction string.""" if img_pil is None: return "Error: Process function received an empty image." try: print(f"Processing PIL Image - Mode: {img_pil.mode}, Size: {img_pil.size}") # Debug # Ensure image is RGB if required by processor/model if img_pil.mode != 'RGB': img_pil = img_pil.convert('RGB') print("Converted image to RGB.") # Debug # Process image with Hugging Face processor inputs = processor(images=img_pil, return_tensors="pt") # Model prediction outputs = model(**inputs) logits = outputs.logits predicted_class_index = np.argmax(logits.detach().numpy(), axis=1)[0] predicted_emotion = labels[predicted_class_index] print(f"Predicted Emotion Index: {predicted_class_index}, Emotion: {predicted_emotion}") # Debug # Get quote quote, author = get_quotes(predicted_class_index, authors=selected_authors) return f"Predicted Emotion: {predicted_emotion.capitalize()}\n\n\"{quote}\"\n\n— {author}" except Exception as e: import traceback print(f"Error in process_pil_image_and_get_quote: {e}\n{traceback.format_exc()}") return f"An error occurred during image processing: {e}" # --- API FUNCTION: Handles Base64 input --- def main_fn_api(img_base64_string, selected_authors): """ API endpoint function. Expects a base64 data URI string for the image. Decodes it and calls the core processing function. """ print(">>> API endpoint main_fn_api called.") # Debug if not isinstance(img_base64_string, str) or not img_base64_string.startswith('data:image'): print(f"Error: Invalid base64 string received: {str(img_base64_string)[:100]}...") # Debug return "Error: Invalid image data received via API." try: # 1. Split header (e.g., "data:image/png;base64,") and encoded data try: header, encoded_data = img_base64_string.split(',', 1) except ValueError: print("Error: Base64 string does not contain expected comma separator.") return "Error: Malformed Base64 image data." # 2. Decode Base64 data print(f"Attempting to decode base64 data (starts with {encoded_data[:10]}...).") # Debug decoded_bytes = base64.b64decode(encoded_data) print(f"Decoded {len(decoded_bytes)} bytes.") # Debug # 3. Create PIL Image from decoded bytes img_pil = Image.open(io.BytesIO(decoded_bytes)) print("Successfully created PIL Image from decoded bytes.") # Debug # 4. Call the core processing function with the PIL image return process_pil_image_and_get_quote(img_pil, selected_authors) except base64.binascii.Error as e: print(f"Error decoding Base64: {e}") # Debug return "Error: Invalid Base64 encoding." except Exception as e: import traceback print(f"Error in main_fn_api: {e}\n{traceback.format_exc()}") return f"An error occurred during API image processing: {e}" # --- Other helper functions (Keep as is) --- def get_random_quote(authors): quote, author = get_quotes(index=None, authors=authors) return f"\"{quote}\"\n\n— {author}" def select_all_authors(): return author_list # --- End other helper functions --- # --- Build Gradio Interface --- with gr.Blocks() as demo: gr.Markdown("## Emotion Detection from Image & Quote") with gr.Row(): # This Image component is primarily for the Web UI interaction image_input_ui = gr.Image(type="pil", label="Upload Image (for Web UI)") output_text = gr.Textbox(label="Prediction & Quote", lines=5) # Increased lines with gr.Row(): author_selector = gr.Dropdown( choices=author_list, multiselect=True, label="Select Author(s)", info="Optional: filter quotes by author" ) with gr.Row(): select_all_button = gr.Button("Select All Authors") with gr.Row(): # This button triggers the processing for images uploaded via the UI submit_button_ui = gr.Button("Detect Emotion & Get Quote (from UI Upload)") random_quote_button = gr.Button("Get a Random Quote") # --- Event Handlers --- # 1. UI Submit Button: Calls the core logic function directly with the PIL image from gr.Image submit_button_ui.click( process_pil_image_and_get_quote, # Target function expects PIL image inputs=[image_input_ui, author_selector], # Inputs from UI components outputs=output_text # No api_name here means this specific button click isn't easily callable via API ) # 2. API Endpoint Registration: # We need to expose `main_fn_api` so the Android app can call it. # We can link it to a component's event and give it the correct `api_name`. # Using a "hidden" Textbox's `change` event is a way to register it without a visible trigger. api_trigger_input = gr.Textbox(visible=False) # Hidden component to hook the API call api_trigger_input.change( # Event doesn't really matter, just need to link the function fn=main_fn_api, # The function that expects BASE64 STRING inputs=[api_trigger_input, author_selector], # Inputs match main_fn_api signature outputs=output_text, # API calls need an output defined api_name="main_fn" # <--- CRITICAL: This name matches the Android API URL target # (.../gradio_api/call/main_fn) ) # 3. Random Quote Button (Keep as is) random_quote_button.click( get_random_quote, inputs=author_selector, outputs=output_text, api_name="get_random_quote" # Keep this if Android uses it too ) # 4. Select All Authors Button (Keep as is) select_all_button.click(select_all_authors, inputs=None, outputs=author_selector) print("Launching Gradio Interface...") demo.launch() # share=True if you need external access temporarily for testing