ipetrousov commited on
Commit
aa40295
·
verified ·
1 Parent(s): c9966f3

Fix app file

Browse files
Files changed (1) hide show
  1. app.py +179 -180
app.py CHANGED
@@ -1,204 +1,203 @@
1
- import spaces
2
- import torch
3
- import torch.nn as nn
4
- from huggingface_hub import hf_hub_download
5
- import librosa
6
- import numpy as np
7
  import gradio as gr
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8
 
 
 
 
 
 
 
9
 
10
- # Vars
11
- repo_name = "ipetrousov/ser_supcon"
12
- device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
13
- weights_dict = "supcon_stage_ft_acc60.pt"
14
-
15
- # Classes
16
- class EmotionFeatureExtractorBase(nn.Module):
17
- """
18
- Core CNN to extract spatial features from Mel-Spectrograms.
19
- Outputs a flat 256-dimensional feature vector.
20
- """
21
-
22
- def __init__(self):
23
- super(EmotionFeatureExtractorBase, self).__init__()
24
-
25
- self.conv1 = nn.Sequential(
26
- nn.Conv2d(in_channels=1, out_channels=32, kernel_size=3, stride=1, padding=1),
27
- nn.BatchNorm2d(32),
28
- nn.ReLU(),
29
- nn.MaxPool2d(kernel_size=2, stride=2)
30
- )
31
-
32
- self.conv2 = nn.Sequential(
33
- nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3, stride=1, padding=1),
34
- nn.BatchNorm2d(64),
35
- nn.ReLU(),
36
- nn.MaxPool2d(kernel_size=2, stride=2)
37
- )
38
-
39
- self.conv3 = nn.Sequential(
40
- nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, stride=1, padding=1),
41
- nn.BatchNorm2d(128),
42
- nn.ReLU(),
43
- nn.MaxPool2d(kernel_size=2, stride=2)
44
- )
45
-
46
- self.conv4 = nn.Sequential(
47
- nn.Conv2d(in_channels=128, out_channels=256, kernel_size=3, stride=1, padding=1),
48
- nn.BatchNorm2d(256),
49
- nn.ReLU(),
50
- nn.MaxPool2d(kernel_size=2, stride=2)
51
- )
52
-
53
- self.global_pool = nn.AdaptiveAvgPool2d((1, 1))
54
- self.flatten = nn.Flatten()
55
-
56
- def forward(self, x):
57
- x = self.conv1(x)
58
- x = self.conv2(x)
59
- x = self.conv3(x)
60
- x = self.conv4(x)
61
- x = self.global_pool(x)
62
- x = self.flatten(x)
63
- return x # Output shape: (batch_size, 256)
64
-
65
-
66
- class SupConStage2(nn.Module):
67
- """ Model B (Stage 2): Frozen Encoder + Linear Classifier. """
68
-
69
- def __init__(self, trained_encoder, num_classes=8):
70
- super(SupConStage2, self).__init__()
71
-
72
- # Pass trained encoder from Stage 1
73
- self.encoder = trained_encoder
74
-
75
- # Freeze encoder weights
76
- for param in self.encoder.parameters():
77
- param.requires_grad = False
78
-
79
- # Linear classifier
80
- self.classifier = nn.Linear(256, num_classes)
81
-
82
- def forward(self, x):
83
- # Disable gradient tracking for the encoder to save memory/compute
84
- with torch.no_grad():
85
- features = self.encoder(x)
86
-
87
- return self.classifier(features)
88
-
89
-
90
- # Function
91
- def preprocess_drive_audio(file_path, target_sample_rate=22050, n_mels=128, noise_std=0.0):
92
- """
93
- Loads audio and generates a normalized Mel-spectrogram using librosa.
94
- Returns a PyTorch tensor shaped (1, 1, Mels, Time).
95
- """
96
-
97
- y, sr = librosa.load(file_path, sr=target_sample_rate, mono=True)
98
-
99
- mel_spectrogram = librosa.feature.melspectrogram(
100
- y=y,
101
- sr=target_sample_rate,
102
- n_fft=512,
103
- hop_length=256,
104
- n_mels=n_mels,
105
- fmax=8000
 
 
 
106
  )
107
 
108
- mel_spectrogram_db = librosa.power_to_db(mel_spectrogram, ref=1.0)
109
-
110
- mean = np.mean(mel_spectrogram_db)
111
- std = np.std(mel_spectrogram_db)
112
- mel_normalized = (mel_spectrogram_db - mean) / (std + 1e-6)
113
-
114
- final_tensor = torch.tensor(mel_normalized, dtype=torch.float32)
115
- final_tensor = final_tensor.unsqueeze(0).unsqueeze(0)
116
-
117
- # Noise Injection Pipeline
118
- if noise_std > 0.0:
119
- noise_matrix = torch.randn_like(final_tensor)
120
- final_tensor = final_tensor + (noise_matrix * noise_std)
121
-
122
- return final_tensor
123
 
 
 
 
 
 
124
 
125
- def run_local_inference(model, file_path, device, noise_std=0.0):
126
- """
127
- Runs a forward pass
128
- """
129
 
130
- emotion_classes = ["Neutral", "Calm", "Happy", "Sad", "Angry", "Fearful", "Disgust", "Surprised"]
 
 
 
 
 
 
131
 
132
- input_tensor = preprocess_drive_audio(file_path, noise_std=noise_std)
133
- input_tensor = input_tensor.to(device)
 
 
134
 
135
- model.eval()
136
- with torch.no_grad():
137
- logits = model(input_tensor)
138
-
139
- # Get the top prediction index
140
- predicted_index = torch.argmax(logits, dim=1).item()
141
- top_class = emotion_classes[predicted_index]
142
-
143
- # Apply softmax to convert logits to probabilities between 0 and 1
144
- probabilities = torch.softmax(logits, dim=1).squeeze().tolist()
145
-
146
- # Map the probabilities to their respective class names for Gradio
147
- prob_dict = {emotion_classes[i]: probabilities[i] for i in range(len(emotion_classes))}
148
-
149
- return top_class, prob_dict
150
-
151
-
152
- # Load SOTA SupConFN
153
- blank_encoder = EmotionFeatureExtractorBase()
154
- model = SupConStage2(trained_encoder=blank_encoder, num_classes=8)
155
-
156
- # Download & Load weights
157
- model_path = hf_hub_download(repo_id=repo_name, filename=weights_dict)
158
- model.load_state_dict(torch.load(model_path, map_location=device))
159
- model.eval()
160
- model = model.to(device)
161
-
162
-
163
- # Gradio Interface Setup
164
- @spaces.GPU
165
- def predict_emotion(audio_path, noise_std):
166
- if audio_path is None:
167
- return "Upload an audio file."
168
-
169
- try:
170
- pred = run_local_inference(model, audio_path, device, noise_std=noise_std)
171
- return pred
172
- except Exception as e:
173
- return f"Error processing audio: {str(e)}"
174
 
175
 
176
  description_html = """
177
- Upload a short audio clip and test the architecture's resilience. The slider injects mathematically controlled Gaussian noise into the normalized Mel-spectrogram before classification.
178
 
179
- <div style="display: flex; gap: 10px; margin-top: 15px;">
180
- <a href="https://github.com/gpetrousov/dl_assignment_demokritos" target="_blank">
181
  <img src="https://img.shields.io/badge/GitHub-View_Repository-181717?style=for-the-badge&logo=github" alt="GitHub Repository" />
182
  </a>
183
- <a href="https://docs.google.com/presentation/d/1ShSmXY2OqCEbQQjYiLYCFc8Psu2bioECoFHhliiDEiQ/edit?usp=drive_link" target="_blank">
 
 
 
184
  <img src="https://img.shields.io/badge/Presentation-View_Slides-F4B400?style=for-the-badge&logo=googleslides&logoColor=white" alt="Presentation Slides" />
185
  </a>
186
  </div>
187
  """
188
 
189
- demo = gr.Interface(
190
- fn=predict_emotion,
191
- inputs=[
192
- gr.Audio(type="filepath", label="Upload Audio (.mp3 or .wav)"),
193
- gr.Slider(minimum=0.0, maximum=1.0, step=0.05, value=0.0, label="Gaussian Noise Injection (STD)")
194
- ],
195
- outputs=[
196
- gr.Textbox(label="Final Prediction"),
197
- gr.Label(num_top_classes=8, label="Probability Distribution")
198
- ],
199
- title="Robust Speech Emotion Recognition (SupCon)",
200
- description=description_html,
201
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
202
 
203
  if __name__ == "__main__":
204
  demo.launch()
 
1
+ import cv2
 
 
 
 
 
2
  import gradio as gr
3
+ import numpy as np
4
+ import matplotlib.pyplot as plt
5
+ from PIL import Image
6
+ import torch
7
+ import open_clip
8
+ import chromadb
9
+
10
+ # 1. Device and Model Initialization
11
+ device = "cuda" if torch.cuda.is_available() else "cpu"
12
+ model, _, preprocess = open_clip.create_model_and_transforms(
13
+ "ViT-B-32",
14
+ pretrained="laion2b_s34b_b79k",
15
+ device=device
16
+ )
17
+ tokenizer = open_clip.get_tokenizer("ViT-B-32")
18
 
19
+ # 2. Visual Attribution Heatmap Generation Function
20
+ def generate_visual_attribution(model, preprocess, rgb_image, query_string, grid_size=4):
21
+ """Generates a side-by-side visualization array with a semantic heatmap overlay."""
22
+ h, w, _ = rgb_image.shape
23
+ patch_h, patch_w = h // grid_size, w // grid_size
24
+ heatmap = np.zeros((grid_size, grid_size))
25
 
26
+ # Tokenize and vectorize text query
27
+ text_token = tokenizer([query_string]).to(device)
28
+ with torch.no_grad():
29
+ text_features = model.encode_text(text_token)
30
+ text_features /= text_features.norm(dim=-1, keepdim=True)
31
+
32
+ # Calculate patch-level cosine similarities
33
+ for i in range(grid_size):
34
+ for j in range(grid_size):
35
+ ymin, ymax = i * patch_h, (i + 1) * patch_h
36
+ xmin, xmax = j * patch_w, (j + 1) * patch_w
37
+
38
+ patch = rgb_image[ymin:ymax, xmin:xmax]
39
+ pil_patch = Image.fromarray(patch)
40
+ patch_tensor = preprocess(pil_patch).unsqueeze(0).to(device)
41
+
42
+ with torch.no_grad():
43
+ patch_features = model.encode_image(patch_tensor)
44
+ patch_features /= patch_features.norm(dim=-1, keepdim=True)
45
+ similarity = (patch_features @ text_features.T).item()
46
+ heatmap[i, j] = similarity
47
+
48
+ heatmap_resized = cv2.resize(heatmap, (w, h), interpolation=cv2.INTER_CUBIC)
49
+
50
+ # Render side-by-side figure
51
+ fig, axes = plt.subplots(1, 2, figsize=(12, 6))
52
+ axes[0].imshow(rgb_image)
53
+ axes[0].set_title("Original Matched Frame")
54
+ axes[0].axis("off")
55
+
56
+ axes[1].imshow(rgb_image)
57
+ axes[1].imshow(heatmap_resized, cmap="jet", alpha=0.5)
58
+ axes[1].set_title(f"Visual Attribution Heatmap for:\n\"{query_string}\"")
59
+ axes[1].axis("off")
60
+
61
+ plt.tight_layout()
62
+
63
+ # Convert Matplotlib figure directly to an RGB array for Gradio
64
+ fig.canvas.draw()
65
+ output_image = np.asarray(fig.canvas.buffer_rgba())[:, :, :3]
66
+ plt.close(fig)
67
+
68
+ return output_image
69
+
70
+ # 3. Main Pipeline Execution
71
+ def process_video_and_query(video_path, query_string, sample_rate_sec=1):
72
+ if not video_path:
73
+ return None, "Please upload a video file."
74
+ if not query_string:
75
+ return None, "Please enter a search query."
76
+
77
+ # Initialize fresh in-memory ChromaDB collection per query
78
+ chroma_client = chromadb.Client()
79
+ collection = chroma_client.create_collection(name="temp_video_analysis")
80
+
81
+ # Ingest and sample video frames
82
+ cap = cv2.VideoCapture(video_path)
83
+ fps = cap.get(cv2.CAP_PROP_FPS)
84
+ if fps == 0 or np.isnan(fps):
85
+ fps = 30.0
86
+ frame_interval = int(fps * sample_rate_sec)
87
+
88
+ frame_count = 0
89
+ saved_frames = {}
90
+ embeddings = []
91
+ ids = []
92
+ metadatas = []
93
+
94
+ while cap.isOpened():
95
+ ret, frame = cap.read()
96
+ if not ret:
97
+ break
98
+
99
+ if frame_count % frame_interval == 0:
100
+ timestamp = round(frame_count / fps, 2)
101
+ rgb_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
102
+ saved_frames[str(timestamp)] = rgb_frame
103
+
104
+ pil_img = Image.fromarray(rgb_frame)
105
+ img_tensor = preprocess(pil_img).unsqueeze(0).to(device)
106
+
107
+ with torch.no_grad():
108
+ img_embedding = model.encode_image(img_tensor).flatten().tolist()
109
+
110
+ embeddings.append(img_embedding)
111
+ ids.append(f"frame_{timestamp}")
112
+ metadatas.append({"timestamp": str(timestamp)})
113
+
114
+ frame_count += 1
115
+ cap.release()
116
+
117
+ if not embeddings:
118
+ return None, "Error: Could not extract frames from the video."
119
+
120
+ # Batch add to ChromaDB
121
+ collection.add(
122
+ embeddings=embeddings,
123
+ ids=ids,
124
+ metadatas=metadatas
125
  )
126
 
127
+ # Vectorize text query
128
+ text_token = tokenizer([query_string]).to(device)
129
+ with torch.no_grad():
130
+ query_embedding = model.encode_text(text_token).flatten().tolist()
 
 
 
 
 
 
 
 
 
 
 
131
 
132
+ # Retrieve nearest neighbor
133
+ results = collection.query(
134
+ query_embeddings=[query_embedding],
135
+ n_results=1
136
+ )
137
 
138
+ best_timestamp = results["metadatas"][0][0]["timestamp"]
139
+ distance = results["distances"][0][0]
140
+ matched_frame = saved_frames[best_timestamp]
 
141
 
142
+ # Generate side-by-side attribution image
143
+ side_by_side_output = generate_visual_attribution(
144
+ model=model,
145
+ preprocess=preprocess,
146
+ rgb_image=matched_frame,
147
+ query_string=query_string
148
+ )
149
 
150
+ status_message = (
151
+ f"Matched Timestamp: {best_timestamp}s\n"
152
+ f"ChromaDB Squared Euclidean Distance: {distance:.4f}"
153
+ )
154
 
155
+ return side_by_side_output, status_message
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
156
 
157
 
158
  description_html = """
159
+ Upload surveillance footage and enter a natural language safety policy to retrieve candidate violation frames via Open-CLIP embeddings and ChromaDB nearest-neighbor search, complete with patch-level visual attribution heatmaps.
160
 
161
+ <div style="display: flex; gap: 10px; margin-top: 15px; flex-wrap: wrap;">
162
+ <a href="https://github.com/gpetrousov/multimodal_ai_assignment_demokritos" target="_blank">
163
  <img src="https://img.shields.io/badge/GitHub-View_Repository-181717?style=for-the-badge&logo=github" alt="GitHub Repository" />
164
  </a>
165
+ <a href="https://github.com/gpetrousov/multimodal_ai_assignment_demokritos/blob/master/report_src/report.pdf" target="_blank">
166
+ <img src="https://img.shields.io/badge/Report-View_PDF-E50914?style=for-the-badge&logo=adobeacrobatreader&logoColor=white" alt="View PDF Report" />
167
+ </a>
168
+ <a href="https://docs.google.com/presentation/d/1DW28UWfK28nR-G_rB0mqrZmNZY3uNZMmDJpCxxdA-fk/edit?usp=sharing" target="_blank">
169
  <img src="https://img.shields.io/badge/Presentation-View_Slides-F4B400?style=for-the-badge&logo=googleslides&logoColor=white" alt="Presentation Slides" />
170
  </a>
171
  </div>
172
  """
173
 
174
+ # 4. Gradio Interface Layout
175
+ with gr.Blocks(title="Zero-Shot Video Safety Auditor") as demo:
176
+ gr.Markdown("# Zero-Shot Video Safety Policy Auditor")
177
+ gr.Markdown(
178
+ "Upload surveillance footage and enter a natural language safety policy "
179
+ "to retrieve candidate violation frames with visual attribution heatmaps."
180
+ )
181
+ gr.HTML(description_html)
182
+
183
+ with gr.Row():
184
+ with gr.Column():
185
+ video_input = gr.Video(label="Input Surveillance Video")
186
+ query_input = gr.Textbox(
187
+ label="Safety Policy Query",
188
+ placeholder="e.g., person operating a yellow construction machine"
189
+ )
190
+ submit_btn = gr.Button("Analyze & Retrieve", variant="primary")
191
+
192
+ with gr.Column():
193
+ image_output = gr.Image(label="Side-by-Side Retrieval & Attribution Heatmap")
194
+ details_output = gr.Textbox(label="Retrieval Metadata", interactive=False)
195
+
196
+ submit_btn.click(
197
+ fn=process_video_and_query,
198
+ inputs=[video_input, query_input],
199
+ outputs=[image_output, details_output]
200
+ )
201
 
202
  if __name__ == "__main__":
203
  demo.launch()