ApurvaKondekar's picture
new app.py
4bfa608 verified
Raw
History Blame
14.7 kB
# =========================================================
# PROFESSIONAL MULTIMODAL EMOTION RECOGNITION UI
# Sky Blue + White Theme
# =========================================================
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"]
gpu_status = (
"GPU Enabled"
if torch.cuda.is_available()
else "Running on CPU"
)
# =========================================================
# LOAD PROCESSORS
# =========================================================
processor = Wav2Vec2Processor.from_pretrained(
"facebook/wav2vec2-base-960h"
)
tokenizer = AutoTokenizer.from_pretrained(
"bert-base-uncased"
)
# =========================================================
# MODEL ARCHITECTURE
# =========================================================
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 = nn.GELU()
self.act2 = 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)
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)
fused = self.hbf(
a_pool,
t_pool,
v_pool
)
logits = self.fc(fused)
return logits
# =========================================================
# LOAD MODEL
# =========================================================
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 successfully")
except Exception as e:
print(f"Model loading failed: {e}")
# =========================================================
# VIDEO PROCESSING
# =========================================================
def extract_video_frames(
video_path,
max_frames=8,
resize=(224, 224)
):
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
return None
total_frames = int(
cap.get(cv2.CAP_PROP_FRAME_COUNT)
)
indices = np.linspace(
0,
total_frames - 1,
max_frames,
dtype=int
)
frames = []
for idx in indices:
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
ret, frame = cap.read()
if not ret:
continue
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])
return frames
# =========================================================
# AUDIO EXTRACTION
# =========================================================
def extract_audio_from_video(video_path):
audio_path = tempfile.NamedTemporaryFile(
delete=False,
suffix=".wav"
).name
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
# =========================================================
# WHISPER
# =========================================================
whisper_model = whisper.load_model("base")
def transcribe_audio(audio_path):
result = whisper_model.transcribe(audio_path)
return result["text"].strip()
# =========================================================
# PREPROCESSING
# =========================================================
def preprocess_inputs(
audio_path,
text,
video_path
):
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)
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
)
# =========================================================
# PREDICTION
# =========================================================
def predict_emotion(video_file):
if video_file is None:
return (
"Please upload a video.",
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():
logits = model(
audio,
audio_mask,
text_ids,
text_mask,
video
)
probs = torch.softmax(
logits,
dim=1
)[0].cpu().numpy()
result = {
LABELS[i]: float(probs[i])
for i in range(len(LABELS))
}
predicted_emotion = LABELS[
probs.argmax()
]
confidence = float(probs.max())
result_text = f"""
## Predicted Emotion
### {predicted_emotion.upper()}
Confidence Score: {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,
""
)
# =========================================================
# CLEAN PROFESSIONAL CSS
# =========================================================
custom_css = """
body {
background: #eaf6ff;
}
.gradio-container {
max-width: 1250px !important;
margin: auto;
padding-top: 20px;
}
.main-title {
text-align: center;
font-size: 48px;
font-weight: 800;
color: #0f4c81;
margin-bottom: 10px;
}
.subtitle {
text-align: center;
font-size: 20px;
color: #3b82b6;
margin-bottom: 30px;
}
.section-box {
background: white;
border-radius: 16px;
padding: 20px;
box-shadow: 0px 4px 12px rgba(0,0,0,0.08);
}
.footer {
text-align: center;
color: #4b5563;
margin-top: 30px;
font-size: 14px;
}
.gr-button {
background: #38bdf8 !important;
border: none !important;
color: white !important;
font-weight: 600 !important;
}
.gr-button:hover {
background: #0ea5e9 !important;
}
h1, h2, h3, h4 {
color: #0f4c81 !important;
}
label {
color: #0f4c81 !important;
font-weight: 600 !important;
}
"""
# =========================================================
# UI
# =========================================================
with gr.Blocks(
title="Multimodal Emotion Recognition",
theme=gr.themes.Soft(),
css=custom_css
) as demo:
# HEADER
gr.HTML("""
<div class="main-title">
Multimodal Emotion Recognition
</div>
<div class="subtitle">
AI-based Emotion Detection using Audio, Text and Video Fusion
</div>
""")
# STATUS
gr.Markdown(
f"### System Status: {gpu_status}"
)
# INFO PANELS
with gr.Row():
with gr.Column():
with gr.Group():
gr.Markdown("""
### Modalities Used
- Audio Analysis
- Speech Transcription
- Facial Expression Analysis
""")
with gr.Column():
with gr.Group():
gr.Markdown("""
### Models Used
- Wav2Vec2
- BERT
- ResNet18
- Whisper
""")
gr.Markdown("---")
# MAIN SECTION
with gr.Row(equal_height=True):
# LEFT
with gr.Column(scale=1):
gr.Markdown("## Upload Video")
video_input = gr.Video(
label="Input Video",
height=400
)
predict_btn = gr.Button(
"Analyze Emotion",
variant="primary",
size="lg"
)
# RIGHT
with gr.Column(scale=1):
gr.Markdown("## Results")
result_text = gr.Markdown(
value="Upload a video and click Analyze Emotion."
)
result_output = gr.Label(
label="Emotion Probabilities",
num_top_classes=4
)
transcription_output = gr.Textbox(
label="Transcribed Text",
lines=6,
interactive=False
)
gr.Markdown("---")
# ABOUT SECTION
with gr.Accordion(
"About the Model",
open=False
):
gr.Markdown("""
This multimodal system combines:
- Audio features using Wav2Vec2
- Text understanding using BERT
- Video feature extraction using ResNet18
- Hybrid Fusion Block for final prediction
Supported emotions:
- Angry
- Happy
- Neutral
- Sad
""")
# FOOTER
gr.HTML("""
<div class="footer">
Built using PyTorch, Transformers, Whisper and Gradio
</div>
""")
# BUTTON ACTION
predict_btn.click(
fn=predict_emotion,
inputs=[video_input],
outputs=[
result_text,
result_output,
transcription_output
],
show_progress=True
)
# =========================================================
# LAUNCH
# =========================================================
if __name__ == "__main__":
demo.queue()
demo.launch()