Try_yolo / app.py
krisha06's picture
Update app.py
52a6170 verified
Raw
History Blame Contribute Delete
3.65 kB
import streamlit as st
from ultralytics import YOLO
from ultralytics.nn.tasks import SegmentationModel
import torch
import cv2
import numpy as np
import tempfile
import os
# Fix for PyTorch 2.6
torch.serialization.add_safe_globals([SegmentationModel])
# Load YOLOv8 segmentation model on CPU
model = YOLO("yolov8n-seg.pt").to("cpu")
st.title("πŸŽ₯ Object Segmentation on Uploaded Video")
uploaded_video = st.file_uploader("Upload a video", type=["mp4", "avi", "mov"])
if uploaded_video:
# Save video temporarily
tfile = tempfile.NamedTemporaryFile(delete=False)
tfile.write(uploaded_video.read())
video_path = tfile.name
st.video(video_path)
st.markdown("### ⏳ Processing video... Please wait.")
cap = cv2.VideoCapture(video_path)
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
fps = int(cap.get(cv2.CAP_PROP_FPS))
output_path = os.path.join(tempfile.gettempdir(), "output_segmentation.mp4")
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
out = cv2.VideoWriter(output_path, fourcc, fps, (width, height))
frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
progress_bar = st.progress(0)
# Generate class color map
np.random.seed(42)
colors = {i: tuple(np.random.randint(0, 256, 3).tolist()) for i in range(80)}
frame_index = 0
while cap.isOpened():
ret, frame = cap.read()
if not ret:
break
try:
results = model.predict(frame, conf=0.3, iou=0.5)
except Exception as e:
st.error(f"Prediction failed on frame {frame_index}: {e}")
break
if results[0].masks is not None:
masks = results[0].masks.data.cpu().numpy()
class_ids = results[0].boxes.cls.cpu().numpy().astype(int)
boxes = results[0].boxes.xyxy.cpu().numpy()
names = results[0].names
for mask, class_id, box in zip(masks, class_ids, boxes):
label = names[class_id]
color = colors.get(class_id, (0, 255, 0))
# Resize and apply mask
resized_mask = cv2.resize(mask, (frame.shape[1], frame.shape[0]))
mask_bool = resized_mask > 0.5
# Overlay
colored_mask = np.zeros_like(frame, dtype=np.uint8)
colored_mask[mask_bool] = color
frame = cv2.addWeighted(frame, 1.0, colored_mask, 0.5, 0)
# Label
x1, y1, x2, y2 = box.astype(int)
font = cv2.FONT_HERSHEY_SIMPLEX
font_scale = 0.6
thickness = 1
label_text = label
(text_width, text_height), _ = cv2.getTextSize(label_text, font, font_scale, thickness)
text_x = x1
text_y = y1 - 10 if y1 - 10 > 10 else y1 + text_height + 10
# Draw label background
cv2.rectangle(frame, (text_x - 2, text_y - text_height - 4),
(text_x + text_width + 2, text_y + 4), (0, 0, 0), -1)
cv2.putText(frame, label_text, (text_x, text_y),
font, font_scale, (255, 255, 255), thickness=1, lineType=cv2.LINE_AA)
out.write(frame)
frame_index += 1
progress_bar.progress(min(frame_index / frame_count, 1.0))
cap.release()
out.release()
progress_bar.empty()
st.success("βœ… Video processing completed!")
with open(output_path, "rb") as f:
st.download_button("πŸ“₯ Download Segmented Video", f, file_name="segmented_output.mp4", mime="video/mp4")