File size: 5,504 Bytes
a818fda
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
from abc import abstractmethod
from typing import List, Optional
import cv2
import subprocess
import numpy as np 

def putText(img, text: str, position, 
               text_font: int=0, text_scale: int=1,
               bg_color=(255,255,255),
               text_color=(255,0,255),
               bg_thickness=8, 
               text_thickness=1, 
               lineType=cv2.LINE_AA):
    """ Function to put text on image.
    
    Args:
        img (_type_): 
        text (str): _description_
        position (_type_): Top-left position of text.
        text_font (int, optional): font size of text. Defaults to 0.
        text_scale (int, optional): text scale. Defaults to 1.
        bg_color (tuple, optional): text background color. Defaults to (255,255,255).
        text_color (tuple, optional): text foreground color. Defaults to (255,0,255).
        bg_thickness (int, optional): text background thickness. Defaults to 8.
        text_thickness (int, optional): text foreground thickness. Defaults to 1.
        lineType (_type_, optional): line type. Defaults to cv2.LINE_AA.

    Returns:
        _type_: _description_
    """
    img = cv2.putText(img, text, position, text_font, text_scale, bg_color, thickness=bg_thickness, lineType=lineType)
    img = cv2.putText(img, text, position, text_font, text_scale, text_color, thickness=text_thickness, lineType=lineType)
    return img 


class BaseVisualizer():

    def __init__(self, class_names: Optional[List[str]], fps: int=-1, min_width: int=-1):
        """ Visualizer class for visualization (track_results + count_results).

        Args:
            class_map_ids (Dict): class mapping dictionary to map model's class to original class. Eg {0: 1, 1: 0, 2: 2, 3: 3} mean we swap class ID between 0 and 1.
            fps (int): FPS for output video. If fps = -1, it will have same fps as input video.
            min_width (int): minimum width for output video (height will be scaled to keep aspect ratio as input video). If min_width = -1, it will have same resolution as input video.
        """

        self.fps = fps
        self.min_width = min_width
        self.class_names = class_names
    
    def init_writer(self, input_video_info: List[int], output_path: str):
        """ Init video writer for write visualized frame to output video.

        Args:
            input_video_info (List[int]): It is a list that includes 4 elements of input video information (fps, width, height, num_frames).
            output_path (str): Path to save output video. 
        """
        if (self.fps == -1):
            self.fps = input_video_info[0]
        self.width, self.height = input_video_info[1], input_video_info[2]
        if (self.min_width > 0):
            out_width = min(self.min_width, self.width)
            self.height = (self.height * out_width)//self.width
            self.width = out_width
        self.output_path = output_path
        self.writer = cv2.VideoWriter(self.output_path, cv2.VideoWriter_fourcc(*"mp4v"), int(self.fps), (self.width, self.height))
        
    @staticmethod
    def get_color(idx):
        idx = idx * 3
        color = ((37 * idx) % 255, (17 * idx) % 255, (29 * idx) % 255)

        return color

    @staticmethod
    def draw_dash_line(img,pt1,pt2,color,thickness=1,style='dotted',gap=20):
        dist =((pt1[0]-pt2[0])**2+(pt1[1]-pt2[1])**2)**.5
        pts= []
        for i in  np.arange(0,dist,gap):
            r=i/dist
            x=int((pt1[0]*(1-r)+pt2[0]*r)+.5)
            y=int((pt1[1]*(1-r)+pt2[1]*r)+.5)
            p = (x,y)
            pts.append(p)
        if len(pts) ==0: 
            return
        if style=='dotted':
            for p in pts:
                cv2.circle(img,p,thickness,color,-1)
        else:
            s=pts[0]
            e=pts[0]
            i=0
            for p in pts:
                s=e
                e=p
                if i%2==1:
                    cv2.line(img,s,e,color,thickness)
                i+=1

    @staticmethod
    def draw_dash_poly(img,pts,color,thickness=1,style='dotted',gap=20):
        """ draw a polygon with dash line.

        Args:
            img (_type_): input image.
            pts (_type_): _description_
            color (_type_): _description_
            thickness (int, optional): _description_. Defaults to 1.
            style (str, optional): _description_. Defaults to 'dotted'.
            gap (int, optional): _description_. Defaults to 20.

        Returns:
            _type_: _description_
        """
        s=pts[0]
        e=pts[0]
        pts.append(pts.pop(0))
        for p in pts:
            s=e
            e=p
            BaseVisualizer.draw_dash_line(img,s,e,color,thickness,style,gap=gap)
        return img

    @staticmethod
    def draw_dash_rect(img,pt1,pt2,color,thickness=1,style='dotted',gap=10):
        pts = [pt1,(pt2[0],pt1[1]),pt2,(pt1[0],pt2[1])] 
        return BaseVisualizer.draw_dash_poly(img,pts,color,thickness,style,gap=gap)   
    
    def close(self):
        """ Function to release video writer. It should be called after finish visualization for all input frames.
        """
        self.writer.release()
    
    def convert(self):
        subprocess.run(f"ffmpeg -y -loglevel quiet -stats -i {self.output_path} -c:v libx264 {self.output_path}".split())

    @abstractmethod
    def visualize(self, *args,**kwargs):
        """ Each project should implement this function to visualize a frame.

        """
        raise NotImplementedError