Spaces:
Running on Zero
Running on Zero
File size: 7,377 Bytes
672d0ac aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 672d0ac aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 aa40295 c9966f3 672d0ac | 1 2 3 4 5 6 7 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 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 | import spaces
import cv2
import gradio as gr
import numpy as np
import matplotlib.pyplot as plt
from PIL import Image
import torch
import open_clip
import chromadb
# 1. Device and Model Initialization
device = "cuda" if torch.cuda.is_available() else "cpu"
model, _, preprocess = open_clip.create_model_and_transforms(
"ViT-B-32",
pretrained="laion2b_s34b_b79k",
device=device
)
tokenizer = open_clip.get_tokenizer("ViT-B-32")
# 2. Visual Attribution Heatmap Generation Function
def generate_visual_attribution(model, preprocess, rgb_image, query_string, grid_size=4):
"""Generates a side-by-side visualization array with a semantic heatmap overlay."""
h, w, _ = rgb_image.shape
patch_h, patch_w = h // grid_size, w // grid_size
heatmap = np.zeros((grid_size, grid_size))
# Tokenize and vectorize text query
text_token = tokenizer([query_string]).to(device)
with torch.no_grad():
text_features = model.encode_text(text_token)
text_features /= text_features.norm(dim=-1, keepdim=True)
# Calculate patch-level cosine similarities
for i in range(grid_size):
for j in range(grid_size):
ymin, ymax = i * patch_h, (i + 1) * patch_h
xmin, xmax = j * patch_w, (j + 1) * patch_w
patch = rgb_image[ymin:ymax, xmin:xmax]
pil_patch = Image.fromarray(patch)
patch_tensor = preprocess(pil_patch).unsqueeze(0).to(device)
with torch.no_grad():
patch_features = model.encode_image(patch_tensor)
patch_features /= patch_features.norm(dim=-1, keepdim=True)
similarity = (patch_features @ text_features.T).item()
heatmap[i, j] = similarity
heatmap_resized = cv2.resize(heatmap, (w, h), interpolation=cv2.INTER_CUBIC)
# Render side-by-side figure
fig, axes = plt.subplots(1, 2, figsize=(12, 6))
axes[0].imshow(rgb_image)
axes[0].set_title("Original Matched Frame")
axes[0].axis("off")
axes[1].imshow(rgb_image)
axes[1].imshow(heatmap_resized, cmap="jet", alpha=0.5)
axes[1].set_title(f"Visual Attribution Heatmap for:\n\"{query_string}\"")
axes[1].axis("off")
plt.tight_layout()
# Convert Matplotlib figure directly to an RGB array for Gradio
fig.canvas.draw()
output_image = np.asarray(fig.canvas.buffer_rgba())[:, :, :3]
plt.close(fig)
return output_image
# 3. Main Pipeline Execution
@spaces.GPU
def process_video_and_query(video_path, query_string, sample_rate_sec=1):
if not video_path:
return None, "Please upload a video file."
if not query_string:
return None, "Please enter a search query."
# Initialize fresh in-memory ChromaDB collection per query
chroma_client = chromadb.Client()
collection = chroma_client.create_collection(name="temp_video_analysis")
# Ingest and sample video frames
cap = cv2.VideoCapture(video_path)
fps = cap.get(cv2.CAP_PROP_FPS)
if fps == 0 or np.isnan(fps):
fps = 30.0
frame_interval = int(fps * sample_rate_sec)
frame_count = 0
saved_frames = {}
embeddings = []
ids = []
metadatas = []
while cap.isOpened():
ret, frame = cap.read()
if not ret:
break
if frame_count % frame_interval == 0:
timestamp = round(frame_count / fps, 2)
rgb_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
saved_frames[str(timestamp)] = rgb_frame
pil_img = Image.fromarray(rgb_frame)
img_tensor = preprocess(pil_img).unsqueeze(0).to(device)
with torch.no_grad():
img_embedding = model.encode_image(img_tensor).flatten().tolist()
embeddings.append(img_embedding)
ids.append(f"frame_{timestamp}")
metadatas.append({"timestamp": str(timestamp)})
frame_count += 1
cap.release()
if not embeddings:
return None, "Error: Could not extract frames from the video."
# Batch add to ChromaDB
collection.add(
embeddings=embeddings,
ids=ids,
metadatas=metadatas
)
# Vectorize text query
text_token = tokenizer([query_string]).to(device)
with torch.no_grad():
query_embedding = model.encode_text(text_token).flatten().tolist()
# Retrieve nearest neighbor
results = collection.query(
query_embeddings=[query_embedding],
n_results=1
)
best_timestamp = results["metadatas"][0][0]["timestamp"]
distance = results["distances"][0][0]
matched_frame = saved_frames[best_timestamp]
# Generate side-by-side attribution image
side_by_side_output = generate_visual_attribution(
model=model,
preprocess=preprocess,
rgb_image=matched_frame,
query_string=query_string
)
status_message = (
f"Matched Timestamp: {best_timestamp}s\n"
f"ChromaDB Squared Euclidean Distance: {distance:.4f}"
)
return side_by_side_output, status_message
description_html = """
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.
<div style="display: flex; gap: 10px; margin-top: 15px; flex-wrap: wrap;">
<a href="https://github.com/gpetrousov/multimodal_ai_assignment_demokritos" target="_blank">
<img src="https://img.shields.io/badge/GitHub-View_Repository-181717?style=for-the-badge&logo=github" alt="GitHub Repository" />
</a>
<a href="https://github.com/gpetrousov/multimodal_ai_assignment_demokritos/blob/master/report_src/report.pdf" target="_blank">
<img src="https://img.shields.io/badge/Report-View_PDF-E50914?style=for-the-badge&logo=adobeacrobatreader&logoColor=white" alt="View PDF Report" />
</a>
<a href="https://docs.google.com/presentation/d/1DW28UWfK28nR-G_rB0mqrZmNZY3uNZMmDJpCxxdA-fk/edit?usp=sharing" target="_blank">
<img src="https://img.shields.io/badge/Presentation-View_Slides-F4B400?style=for-the-badge&logo=googleslides&logoColor=white" alt="Presentation Slides" />
</a>
</div>
"""
# 4. Gradio Interface Layout
with gr.Blocks(title="Zero-Shot Video Safety Auditor") as demo:
gr.Markdown("# Zero-Shot Video Safety Policy Auditor")
gr.Markdown(
"Upload surveillance footage and enter a natural language safety policy "
"to retrieve candidate violation frames with visual attribution heatmaps."
)
gr.HTML(description_html)
with gr.Row():
with gr.Column():
video_input = gr.Video(label="Input Surveillance Video")
query_input = gr.Textbox(
label="Safety Policy Query",
placeholder="e.g., person operating a yellow construction machine"
)
submit_btn = gr.Button("Analyze & Retrieve", variant="primary")
with gr.Column():
image_output = gr.Image(label="Side-by-Side Retrieval & Attribution Heatmap")
details_output = gr.Textbox(label="Retrieval Metadata", interactive=False)
submit_btn.click(
fn=process_video_and_query,
inputs=[video_input, query_input],
outputs=[image_output, details_output]
)
if __name__ == "__main__":
demo.launch(ssr_mode=False) |