computer-vision / app.py
raz-1412's picture
Upload 2 files
97c2eaa verified
Raw
History Blame Contribute Delete
2.24 kB
import gradio as gr
import torch
import cv2
import numpy as np
import tempfile
import json
from transformers import DetrImageProcessor, DetrForObjectDetection
from PIL import Image
# Load a pretrained object detection model
model_name = "facebook/detr-resnet-50"
processor = DetrImageProcessor.from_pretrained(model_name)
model = DetrForObjectDetection.from_pretrained(model_name)
# Video processing function
def process_video(video_path):
cap = cv2.VideoCapture(video_path)
frame_rate = cap.get(cv2.CAP_PROP_FPS)
timestamped_objects = []
frame_number = 0
while cap.isOpened():
ret, frame = cap.read()
if not ret:
break
# Process every 10th frame for efficiency
if frame_number % int(frame_rate) == 0:
timestamp = frame_number / frame_rate
pil_image = Image.fromarray(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB))
# Model prediction
inputs = processor(images=pil_image, return_tensors="pt")
with torch.no_grad():
outputs = model(**inputs)
# Extract results
target_sizes = torch.tensor([pil_image.size])
results = processor.post_process_object_detection(outputs, target_sizes=target_sizes, threshold=0.7)[0]
detected_objects = []
for score, label, box in zip(results["scores"], results["labels"], results["boxes"]):
label_name = model.config.id2label[label.item()]
detected_objects.append(label_name)
if detected_objects:
timestamped_objects.append({"timestamp": round(timestamp, 2), "objects": list(set(detected_objects))})
frame_number += 1
cap.release()
return json.dumps(timestamped_objects, indent=2)
# Gradio Interface
iface = gr.Interface(
fn=process_video,
inputs=gr.Video(label="Upload Video"),
outputs=gr.JSON(label="Timestamped Object Detections"),
title="Video Object Detector",
description="Upload a video and get a timestamped list of detected objects using a pretrained DETR model."
)
# Launch the Gradio app
if __name__ == "__main__":
iface.launch()