ApurvaKondekar's picture
new
fd5c3ae verified
Raw
History Blame Contribute Delete
11 kB
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()