Spaces:
Running on Zero
Running on Zero
| 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 | |
| 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) |