| import cv2 |
| from ultralytics import YOLO |
| import numpy as np |
| import PIL |
| import streamlit as st |
| import io |
| import tempfile |
| import moviepy.editor as mpy |
| import random |
|
|
| class ObjectDetectionModel(): |
| def __init__(self): |
| self.model = YOLO("models/yolov8n_openvino_model", task = "detect") |
| self.class_names = self.model.names |
|
|
| with open("utils/game_classes.txt", "r") as file: |
| self.game_classes = [line.strip() for line in file if line.strip()] |
|
|
| def process(self, img): |
| if isinstance(img, np.ndarray): |
| uploaded_img_cv = img |
| else: |
| uploaded_img = PIL.Image.open(img) |
| uploaded_img_cv = np.array(uploaded_img) |
| if uploaded_img_cv.shape[-1] == 4: |
| uploaded_img_cv = cv2.cvtColor(uploaded_img_cv, cv2.COLOR_RGBA2RGB) |
| |
| result = self.model(uploaded_img_cv) |
| img_plot = result[0].plot() |
|
|
| detected_classes = set() |
| for box in result[0].boxes: |
| class_id = int(box.cls[0]) |
| class_name = self.class_names[class_id] |
| detected_classes.add(class_name) |
| detected_objects = f'Objects Detected: {", ".join(detected_classes) if detected_classes else "No objects detected"}' |
|
|
| return img_plot, detected_objects |
| |
| def play_video(self, video_path): |
| uploaded_video = io.BytesIO(video_path.read()) |
| temporary_location = "upload.mp4" |
| |
| with open(temporary_location, "wb") as temp_out: |
| temp_out.write(uploaded_video.read()) |
| temp_out.close() |
| |
| camera = cv2.VideoCapture(temporary_location) |
| fps = camera.get(cv2.CAP_PROP_FPS) |
| temp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.mp4') |
| processed_frames = [] |
|
|
| total_frames = int(camera.get(cv2.CAP_PROP_FRAME_COUNT)) |
| frame_count = 0 |
| progress_bar = st.progress(0) |
| st_frame = st.empty() |
| |
| while(True): |
| ret, frame = camera.read() |
| if not ret: |
| break |
| |
| result = self.model(frame, verbose=False) |
| img_plot = result[0].plot() |
| img_plot_rgb = cv2.cvtColor(img_plot, cv2.COLOR_BGR2RGB) |
| processed_frames.append(img_plot_rgb) |
|
|
| st_frame.image(img_plot, channels = "BGR") |
| frame_count +=1 |
| progress_bar.progress(frame_count/total_frames, text = None) |
| |
| camera.release() |
| clip = mpy.ImageSequenceClip(processed_frames, fps=fps) |
| clip.write_videofile(temp_file.name) |
| |
| st_frame.empty() |
| progress_bar.empty() |
| |
| return temp_file.name |
|
|
| def call_class(self): |
| random_class = random.choice(list(self.game_classes)) |
| |
| return random_class |
|
|