import gradio as gr import torch import torch.nn as nn import numpy as np import librosa import cv2 import re from transformers import Wav2Vec2Processor, Wav2Vec2Model, AutoTokenizer, AutoModel from torchvision import models import tempfile import os from huggingface_hub import hf_hub_download import whisper import subprocess # Configuration DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") SAMPLE_RATE = 16000 TEXT_MAX_LEN = 64 LABELS = ["angry", "happy", "neutral", "sad"] # Load processors processor = Wav2Vec2Processor.from_pretrained("facebook/wav2vec2-base-960h") tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased") # Model Architecture (same as training) class ResNetVideoEncoder(nn.Module): def __init__(self, out_dim=768): super().__init__() base = models.resnet18(pretrained=False) self.backbone = nn.Sequential(*list(base.children())[:-1]) self.proj = nn.Linear(512, out_dim) def forward(self, x): B, C, T, H, W = x.shape feats = [] for t in range(T): ft = self.backbone(x[:, :, t]) feats.append(ft.squeeze(-1).squeeze(-1)) feats = torch.stack(feats, dim=1).mean(1) return self.proj(feats) def mean_pool(x, mask): mask = mask[:, :x.size(1)] mask = mask.unsqueeze(-1).float() return (x * mask).sum(1) / mask.sum(1).clamp(min=1e-6) class HBF(nn.Module): def __init__(self, d=768, n_layers=6): super().__init__() self.proj_a = nn.ModuleList([nn.Linear(d, d) for _ in range(n_layers)]) self.proj_t = nn.ModuleList([nn.Linear(d, d) for _ in range(n_layers)]) self.proj_v = nn.ModuleList([nn.Linear(d, d) for _ in range(n_layers)]) self.fwd1 = nn.ModuleList([nn.Linear(3*d, d) for _ in range(n_layers)]) self.fwd2 = nn.ModuleList([nn.Linear(d, d) for _ in range(n_layers)]) self.drop = nn.Dropout(0.1) self.act1, self.act2 = nn.GELU(), nn.Tanh() self.n = n_layers def forward(self, a, t, v): v_prev = None for i in range(self.n): va = self.act2(self.drop(self.proj_a[i](a))) vt = self.act2(self.drop(self.proj_t[i](t))) vv = self.act2(self.drop(self.proj_v[i](v))) cat = torch.cat([va, vt, vv] if v_prev is None else [va, vt, v_prev], -1) x = self.act1(self.fwd1[i](cat)) v_prev = self.fwd2[i](x) return v_prev class AVVideoModel(nn.Module): def __init__(self, num_classes, n_layers=6): super().__init__() self.a_enc = Wav2Vec2Model.from_pretrained("facebook/wav2vec2-base-960h") self.t_enc = AutoModel.from_pretrained("bert-base-uncased") self.v_enc = ResNetVideoEncoder() self.hbf = HBF(n_layers=n_layers) self.fc = nn.Linear(768, num_classes) self.fc_audio = nn.Linear(768, num_classes) self.fc_text = nn.Linear(768, num_classes) self.fc_video = nn.Linear(768, num_classes) def forward(self, audio, audio_mask, text_ids, text_mask, video): a_out = self.a_enc(audio, attention_mask=audio_mask, return_dict=True) t_out = self.t_enc(input_ids=text_ids, attention_mask=text_mask, return_dict=True) a_pool = mean_pool(a_out.last_hidden_state, audio_mask) t_pool = mean_pool(t_out.last_hidden_state, text_mask) v_pool = self.v_enc(video) a_pool = torch.nan_to_num(a_pool, nan=0.0, posinf=1e4, neginf=-1e4) t_pool = torch.nan_to_num(t_pool, nan=0.0, posinf=1e4, neginf=-1e4) v_pool = torch.nan_to_num(v_pool, nan=0.0, posinf=1e4, neginf=-1e4) a_logits = self.fc_audio(a_pool) t_logits = self.fc_text(t_pool) v_logits = self.fc_video(v_pool) fused = self.hbf(a_pool, t_pool, v_pool) fused = torch.nan_to_num(fused, nan=0.0, posinf=1e4, neginf=-1e4) fused_logits = self.fc(fused) return fused_logits, a_logits, t_logits, v_logits # Load model model = AVVideoModel(num_classes=len(LABELS)).to(DEVICE) # Download model from Hugging Face Model Hub try: model_path = hf_hub_download( repo_id="ApurvaKondekar/emotion_model", # CHANGE THIS filename="model_weights.pth" # YOUR FILE NAME ) model.load_state_dict(torch.load(model_path, map_location=DEVICE)) model.eval() print("✅ Model loaded from Hugging Face") except Exception as e: print(f"❌ Failed to load model: {e}") # Alternative: Create a dummy model for testing # Load trained weights (you'll need to upload this) if os.path.exists("model_weights.pth"): model.load_state_dict(torch.load("model_weights.pth", map_location=DEVICE)) model.eval() print("✅ Model loaded successfully") else: print("⚠️ No model weights found. Using untrained model for demo.") def extract_video_frames(video_path, max_frames=8, resize=(224, 224)): """Extract frames from video file""" cap = cv2.VideoCapture(video_path) if not cap.isOpened(): return None frames = [] while len(frames) < max_frames: ret, frame = cap.read() if not ret: break frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) frame = cv2.resize(frame, resize) frames.append(frame) cap.release() if len(frames) == 0: return None # Pad if needed while len(frames) < max_frames: frames.append(frames[-1]) frames = np.array(frames[:max_frames], dtype=np.uint8) return frames def extract_audio_from_video(video_path): """Extract audio from video file using ffmpeg""" try: audio_path = tempfile.NamedTemporaryFile(delete=False, suffix=".wav").name # Use ffmpeg to extract audio import subprocess command = [ 'ffmpeg', '-i', video_path, '-vn', # No video '-acodec', 'pcm_s16le', # Audio codec '-ar', str(SAMPLE_RATE), # Sample rate '-ac', '1', # Mono '-y', # Overwrite audio_path ] subprocess.run(command, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True) return audio_path except Exception as e: raise ValueError(f"Could not extract audio from video: {str(e)}") def transcribe_audio(audio_path): """Transcribe audio using Whisper""" try: whisper_model = whisper.load_model("base") result = whisper_model.transcribe(audio_path) return result["text"].strip() except Exception as e: raise ValueError(f"Could not transcribe audio: {str(e)}") def preprocess_inputs(audio_path, text, video_path): """Preprocess all three modalities""" # Audio wav, _ = librosa.load(audio_path, sr=SAMPLE_RATE) audio_inputs = processor(wav, sampling_rate=SAMPLE_RATE, return_tensors="pt") audio_values = audio_inputs.input_values.to(DEVICE) audio_mask = torch.ones_like(audio_values).to(DEVICE) # Text text_clean = re.sub(r"[^a-zA-Z0-9\s]", "", text.lower()) text_inputs = tokenizer( text_clean, truncation=True, padding="max_length", max_length=TEXT_MAX_LEN, return_tensors="pt" ) text_ids = text_inputs.input_ids.to(DEVICE) text_mask = text_inputs.attention_mask.to(DEVICE) # Video frames = extract_video_frames(video_path) if frames is None: raise ValueError("Could not extract frames from video") frames_tensor = torch.tensor(frames).permute(0, 3, 1, 2).float() / 255.0 frames_tensor = frames_tensor.unsqueeze(0).permute(0, 2, 1, 3, 4).to(DEVICE) return audio_values, audio_mask, text_ids, text_mask, frames_tensor def predict_emotion(video_file): """Main prediction function - takes only video input""" if video_file is None: return "Please provide a video file", None, "" try: # Extract audio from video audio_path = extract_audio_from_video(video_file) # Transcribe audio transcribed_text = transcribe_audio(audio_path) # Preprocess all modalities audio, audio_mask, text_ids, text_mask, video = preprocess_inputs( audio_path, transcribed_text, video_file ) # Inference with torch.no_grad(): fused_logits, a_logits, t_logits, v_logits = model( audio, audio_mask, text_ids, text_mask, video ) # Get probabilities fused_probs = torch.softmax(fused_logits, dim=1)[0].cpu().numpy() # Format results result = {LABELS[i]: float(fused_probs[i]) for i in range(len(LABELS))} predicted_emotion = LABELS[fused_probs.argmax()] confidence = float(fused_probs.max()) result_text = f"🎯 **Predicted Emotion: {predicted_emotion.upper()}**\n\n**Confidence: {confidence:.2%}**" # Clean up temporary audio file if os.path.exists(audio_path): os.remove(audio_path) return result_text, result, transcribed_text except Exception as e: return f"Error: {str(e)}", None, "" # Gradio Interface with gr.Blocks(title="Multimodal Emotion Recognition", theme=gr.themes.Soft()) as demo: gr.Markdown( """ # 🎭 Multimodal Emotion Recognition This system predicts emotions from video by automatically extracting and analyzing: - 🎤 **Audio** (extracted from video) - 📝 **Text** (transcribed from audio using Whisper) - 🎥 **Video** (visual frames) ### How to use: 1. Upload a video file (MP4, AVI, MOV, etc.) 2. Click "Predict Emotion" 3. The system will automatically extract audio, transcribe speech, and analyze all modalities The model will provide emotion predictions based on all three inputs. """ ) with gr.Row(): with gr.Column(): video_input = gr.Video(label="🎥 Video Input") predict_btn = gr.Button("🚀 Predict Emotion", variant="primary", size="lg") with gr.Column(): result_text = gr.Markdown(label="Result") result_output = gr.Label(label="📊 Prediction Results", num_top_classes=4) transcription_output = gr.Textbox(label="📝 Transcribed Text", lines=3, interactive=False) predict_btn.click( fn=predict_emotion, inputs=[video_input], outputs=[result_text, result_output, transcription_output] ) gr.Markdown( """ --- ### 📌 Notes: - Supported emotions: **Angry, Happy, Neutral, Sad** - Model uses Wav2Vec2 (audio), BERT (text), and ResNet18 (video) - Best results with clear audio, accurate transcripts, and visible faces """ ) if __name__ == "__main__": demo.launch()