| 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 |
|
|
| |
| 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="your-username/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}") |
|
|
| 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 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(audio_file, text_input, video_file): |
| """Main prediction function""" |
| |
| if audio_file is None or video_file is None or not text_input.strip(): |
| return "Please provide all three inputs: audio, text, and video", None, None, None, None |
| |
| try: |
| |
| audio, audio_mask, text_ids, text_mask, video = preprocess_inputs( |
| audio_file, text_input, 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() |
| audio_probs = torch.softmax(a_logits, dim=1)[0].cpu().numpy() |
| text_probs = torch.softmax(t_logits, dim=1)[0].cpu().numpy() |
| video_probs = torch.softmax(v_logits, dim=1)[0].cpu().numpy() |
| |
| |
| fused_result = {LABELS[i]: float(fused_probs[i]) for i in range(len(LABELS))} |
| audio_result = {LABELS[i]: float(audio_probs[i]) for i in range(len(LABELS))} |
| text_result = {LABELS[i]: float(text_probs[i]) for i in range(len(LABELS))} |
| video_result = {LABELS[i]: float(video_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%}**" |
| |
| return result_text, fused_result, audio_result, text_result, video_result |
| |
| except Exception as e: |
| return f"Error: {str(e)}", None, None, None, None |
|
|
| |
| with gr.Blocks(title="Multimodal Emotion Recognition", theme=gr.themes.Soft()) as demo: |
| gr.Markdown( |
| """ |
| # π Multimodal Emotion Recognition |
| |
| This system predicts emotions using **Audio**, **Text**, and **Video** inputs simultaneously. |
| |
| ### How to use: |
| 1. Upload an audio file (WAV, MP3) |
| 2. Enter the transcript or spoken text |
| 3. Upload a video file (MP4, AVI) |
| 4. Click "Predict Emotion" |
| |
| The model will analyze all three modalities and provide predictions. |
| """ |
| ) |
| |
| with gr.Row(): |
| with gr.Column(): |
| audio_input = gr.Audio(type="filepath", label="π€ Audio Input") |
| text_input = gr.Textbox( |
| label="π Text Transcript", |
| placeholder="Enter what was said in the audio/video...", |
| lines=3 |
| ) |
| 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") |
| |
| with gr.Accordion("π Detailed Predictions", open=True): |
| fused_output = gr.Label(label="π Fused Prediction", num_top_classes=4) |
| audio_output = gr.Label(label="π€ Audio-only Prediction", num_top_classes=4) |
| text_output = gr.Label(label="π Text-only Prediction", num_top_classes=4) |
| video_output = gr.Label(label="π₯ Video-only Prediction", num_top_classes=4) |
| |
| predict_btn.click( |
| fn=predict_emotion, |
| inputs=[audio_input, text_input, video_input], |
| outputs=[result_text, fused_output, audio_output, text_output, video_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() |
|
|