opencv/object_tracking_vittrack
4
1# This file is part of OpenCV Zoo project.2# It is subject to the license terms in the LICENSE file found in the same directory.3 4import argparse5 6import numpy as np7import cv2 as cv8 9# Check OpenCV version10opencv_python_version = lambda str_version: tuple(map(int, (str_version.split("."))))11assert opencv_python_version(cv.__version__) >= opencv_python_version("4.10.0"), \12 "Please install latest opencv-python for benchmark: python3 -m pip install --upgrade opencv-python"13 14from vittrack import VitTrack15 16# Valid combinations of backends and targets17backend_target_pairs = [18 [cv.dnn.DNN_BACKEND_OPENCV, cv.dnn.DNN_TARGET_CPU],19 [cv.dnn.DNN_BACKEND_CUDA, cv.dnn.DNN_TARGET_CUDA],20 [cv.dnn.DNN_BACKEND_CUDA, cv.dnn.DNN_TARGET_CUDA_FP16],21 [cv.dnn.DNN_BACKEND_TIMVX, cv.dnn.DNN_TARGET_NPU],22 [cv.dnn.DNN_BACKEND_CANN, cv.dnn.DNN_TARGET_NPU]23]24 25parser = argparse.ArgumentParser(26 description="VIT track opencv API")27parser.add_argument('--input', '-i', type=str,28 help='Usage: Set path to the input video. Omit for using default camera.')29parser.add_argument('--model_path', type=str, default='object_tracking_vittrack_2023sep.onnx',30 help='Usage: Set model path')31parser.add_argument('--backend_target', '-bt', type=int, default=0,32 help='''Choose one of the backend-target pair to run this demo:33 {:d}: (default) OpenCV implementation + CPU,34 {:d}: CUDA + GPU (CUDA),35 {:d}: CUDA + GPU (CUDA FP16),36 {:d}: TIM-VX + NPU,37 {:d}: CANN + NPU38 '''.format(*[x for x in range(len(backend_target_pairs))]))39parser.add_argument('--save', '-s', action='store_true', default=False,40 help='Usage: Specify to save a file with results.')41parser.add_argument('--vis', '-v', action='store_true', default=True,42 help='Usage: Specify to open a new window to show results.')43args = parser.parse_args()44def visualize(image, bbox, score, isLocated, fps=None, box_color=(0, 255, 0),text_color=(0, 255, 0), fontScale = 1, fontSize = 1):45 output = image.copy()46 h, w, _ = output.shape47 48 if fps is not None:49 cv.putText(output, 'FPS: {:.2f}'.format(fps), (0, 30), cv.FONT_HERSHEY_DUPLEX, fontScale, text_color, fontSize)50 51 if isLocated and score >= 0.3:52 # bbox: Tuple of length 453 x, y, w, h = bbox54 cv.rectangle(output, (x, y), (x+w, y+h), box_color, 2)55 cv.putText(output, '{:.2f}'.format(score), (x, y+25), cv.FONT_HERSHEY_DUPLEX, fontScale, text_color, fontSize)56 else:57 text_size, baseline = cv.getTextSize('Target lost!', cv.FONT_HERSHEY_DUPLEX, fontScale, fontSize)58 text_x = int((w - text_size[0]) / 2)59 text_y = int((h - text_size[1]) / 2)60 cv.putText(output, 'Target lost!', (text_x, text_y), cv.FONT_HERSHEY_DUPLEX, fontScale, (0, 0, 255), fontSize)61 62 return output63 64if __name__ == '__main__':65 backend_id = backend_target_pairs[args.backend_target][0]66 target_id = backend_target_pairs[args.backend_target][1]67 68 model = VitTrack(69 model_path=args.model_path,70 backend_id=backend_id,71 target_id=target_id)72 73 # Read from args.input74 _input = 0 if args.input is None else args.input75 video = cv.VideoCapture(_input)76 77 # Select an object78 has_frame, first_frame = video.read()79 if not has_frame:80 print('No frames grabbed!')81 exit()82 first_frame_copy = first_frame.copy()83 cv.putText(first_frame_copy, "1. Drag a bounding box to track.", (0, 25), cv.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0))84 cv.putText(first_frame_copy, "2. Press ENTER to confirm", (0, 50), cv.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0))85 roi = cv.selectROI('VitTrack Demo', first_frame_copy)86 87 if np.all(np.array(roi) == 0):88 print("No ROI is selected! Exiting ...")89 exit()90 else:91 print("Selected ROI: {}".format(roi))92 93 if args.save:94 fps = video.get(cv.CAP_PROP_FPS)95 frame_size = (first_frame.shape[1], first_frame.shape[0])96 output_video = cv.VideoWriter('output.mp4', cv.VideoWriter_fourcc(*'mp4v'), fps, frame_size)97 98 # Init tracker with ROI99 model.init(first_frame, roi)100 101 # Track frame by frame102 tm = cv.TickMeter()103 while cv.waitKey(1) < 0:104 has_frame, frame = video.read()105 if not has_frame:106 print('End of video')107 break108 # Inference109 tm.start()110 isLocated, bbox, score = model.infer(frame)111 tm.stop()112 # Visualize113 frame = visualize(frame, bbox, score, isLocated, fps=tm.getFPS())114 if args.save:115 output_video.write(frame)116 117 if args.vis:118 cv.imshow('VitTrack Demo', frame)119 tm.reset()120 121 if args.save:122 output_video.release()123 124 video.release()125 cv.destroyAllWindows()126 