tti_retrieval / app.py
ipetrousov's picture
Fixed spaces CUDA
672d0ac verified
Raw
History Blame Contribute Delete
7.38 kB
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)