Spaces:
Running on Zero
Running on Zero
Fix app file
Browse files
app.py
CHANGED
|
@@ -1,204 +1,203 @@
|
|
| 1 |
-
import
|
| 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 |
-
#
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
-
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
| 93 |
-
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 105 |
-
|
|
|
|
|
|
|
|
|
|
| 106 |
)
|
| 107 |
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
|
| 111 |
-
|
| 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 |
-
|
| 126 |
-
""
|
| 127 |
-
|
| 128 |
-
"""
|
| 129 |
|
| 130 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 131 |
|
| 132 |
-
|
| 133 |
-
|
|
|
|
|
|
|
| 134 |
|
| 135 |
-
|
| 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
|
| 178 |
|
| 179 |
-
<div style="display: flex; gap: 10px; margin-top: 15px;">
|
| 180 |
-
<a href="https://github.com/gpetrousov/
|
| 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://
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 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()
|