| 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 |
|
|
| |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| SAMPLE_RATE = 16000 |
| TEXT_MAX_LEN = 64 |
| LABELS = ["angry", "happy", "neutral", "sad"] |
|
|
| |
| processor = Wav2Vec2Processor.from_pretrained("facebook/wav2vec2-base-960h") |
| tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased") |
|
|
| |
| 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 |
|
|
| |
|
|
| model = AVVideoModel(num_classes=len(LABELS)).to(DEVICE) |
|
|
| |
| try: |
| model_path = hf_hub_download( |
| repo_id="ApurvaKondekar/emotion_model", |
| filename="model_weights.pth" |
| ) |
| 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}") |
| |
| |
| 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 |
| |
| |
| 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 |
| |
| import subprocess |
| command = [ |
| 'ffmpeg', '-i', video_path, |
| '-vn', |
| '-acodec', 'pcm_s16le', |
| '-ar', str(SAMPLE_RATE), |
| '-ac', '1', |
| '-y', |
| 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""" |
| |
| |
| 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_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) |
| |
| |
| 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: |
| |
| audio_path = extract_audio_from_video(video_file) |
| |
| |
| transcribed_text = transcribe_audio(audio_path) |
| |
| |
| audio, audio_mask, text_ids, text_mask, video = preprocess_inputs( |
| audio_path, transcribed_text, video_file |
| ) |
| |
| |
| with torch.no_grad(): |
| fused_logits, a_logits, t_logits, v_logits = model( |
| audio, audio_mask, text_ids, text_mask, video |
| ) |
| |
| |
| fused_probs = torch.softmax(fused_logits, dim=1)[0].cpu().numpy() |
| |
| |
| 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%}**" |
| |
| |
| 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, "" |
|
|
| |
| 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() |