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)