CoolFace
Modelpublic

opencv/object_tracking_vittrack

sourceHugging Faceupdated 1y agoView on Hugging Face
4likes
demo.py126 linesDownload Raw Back to root
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