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.
GitHub Repository View PDF Report Presentation Slides
""" # 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)