| |
| |
| |
| |
|
|
| 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"] |
|
|
| gpu_status = ( |
| "GPU Enabled" |
| if torch.cuda.is_available() |
| else "Running on CPU" |
| ) |
|
|
| |
| |
| |
|
|
| 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 = 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 |
|
|
| |
| |
| |
|
|
| 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}") |
|
|
| |
| |
| |
|
|
| 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 |
|
|
| |
| |
| |
|
|
| 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_model = whisper.load_model("base") |
|
|
| def transcribe_audio(audio_path): |
|
|
| result = whisper_model.transcribe(audio_path) |
|
|
| return result["text"].strip() |
|
|
| |
| |
| |
|
|
| 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 |
| ) |
|
|
| |
| |
| |
|
|
| 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, |
| "" |
| ) |
|
|
| |
| |
| |
|
|
| 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; |
| } |
| """ |
|
|
| |
| |
| |
|
|
| with gr.Blocks( |
| title="Multimodal Emotion Recognition", |
| theme=gr.themes.Soft(), |
| css=custom_css |
| ) as demo: |
|
|
| |
|
|
| gr.HTML(""" |
| <div class="main-title"> |
| Multimodal Emotion Recognition |
| </div> |
| |
| <div class="subtitle"> |
| AI-based Emotion Detection using Audio, Text and Video Fusion |
| </div> |
| """) |
|
|
| |
|
|
| gr.Markdown( |
| f"### System Status: {gpu_status}" |
| ) |
|
|
| |
|
|
| 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("---") |
|
|
| |
|
|
| with gr.Row(equal_height=True): |
|
|
| |
|
|
| 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" |
| ) |
|
|
| |
|
|
| 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("---") |
|
|
| |
|
|
| 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 |
| """) |
|
|
| |
|
|
| gr.HTML(""" |
| <div class="footer"> |
| Built using PyTorch, Transformers, Whisper and Gradio |
| </div> |
| """) |
|
|
| |
|
|
| predict_btn.click( |
| fn=predict_emotion, |
| inputs=[video_input], |
| outputs=[ |
| result_text, |
| result_output, |
| transcription_output |
| ], |
| show_progress=True |
| ) |
|
|
| |
| |
| |
|
|
| if __name__ == "__main__": |
|
|
| demo.queue() |
|
|
| demo.launch() |