Spaces:
Build error
Build error
| 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() | |