diff --git a/ocv_corner_tracker.py b/ocv_corner_tracker.py index 3356607..077b5e4 100644 --- a/ocv_corner_tracker.py +++ b/ocv_corner_tracker.py @@ -99,9 +99,16 @@ class Corner: return self._path +class CornerTrackerParams: + def __init__(self, _scale: float = 1.0, _pre_track: bool = False, _show_path: bool = False): + self.scale = _scale + self.pre_track = _pre_track + self.show_path = _show_path + + class CornerTracker: - def __init__(self, _do_track, color=(0, 255, 0), name: str = 'CornerTracker'): - self.do_track = _do_track + def __init__(self, _params: CornerTrackerParams, color=(0, 255, 0), name: str = 'CornerTracker'): + self.params = _params self.color = color self.name = name self.matching_tpl_bb = None @@ -208,7 +215,7 @@ class CornerTracker: return False # Initialize tracker with first frame and bounding box - if self.do_track: + if self.params.pre_track: self.tracker = cv2.TrackerKCF.create() self.tracker.init(_image, self.tracking_bb) @@ -230,7 +237,7 @@ class CornerTracker: if not _ok: return None - _image_local = image_crop(_image.copy(), self.tracking_bb).copy() + _image_local = image_crop(_image, self.tracking_bb) CornerTracker.mask_apply(_image_local, self.tracking_mask) _image_processed_local = CornerTracker.image_process(_image_local) @@ -258,17 +265,18 @@ class CornerTracker: corners_refined = np.array(corners_refined) - # Create path from global refined corners - _i = 0 - for corner in corners_refined: - _ct = self.corner_matcher_list[_i] - _ct.path_add(corner) - _i += 1 + if self.params.show_path: + # Create path from global refined corners + _i = 0 + for corner in corners_refined: + _ct = self.corner_matcher_list[_i] + _ct.path_add(corner) + _i += 1 - # draw path - for _ct in self.corner_matcher_list: - for p in _ct.path: - cv2.line(_image_anno, bbox_round(p['from']), bbox_round(p['to']), COLOR_TRACK, 1) + # draw path + for _ct in self.corner_matcher_list: + for p in _ct.path: + cv2.line(_image_anno, bbox_round(p['from']), bbox_round(p['to']), COLOR_TRACK, 1) distances = [] for _i in range(0, corners_refined.shape[0]): @@ -319,18 +327,14 @@ if __name__ == '__main__': # Parse command line args parser = argparse.ArgumentParser(description="Gearbox Tracker") parser.add_argument("filename") - parser.add_argument("--scale") + parser.add_argument("--scale", default=1.0) parser.add_argument("--track", action="store_true") + parser.add_argument("--path", action="store_true") args = parser.parse_args() video = cv2.VideoCapture(args.filename) - # Parse scale - scale = 1.0 - if args.scale is not None: - scale = np.float32(args.scale) - - # Parse track - do_track = args.track + # Fill CornerTrackerParams from args + ct_params = CornerTrackerParams(_scale=np.float32(args.scale), _pre_track=args.track, _show_path=args.path) # Let's go colors = [(255, 0, 0), (0, 255, 0), (0, 0, 255), (255, 0, 255), (0, 255, 255), (255, 255, 255)] @@ -343,7 +347,7 @@ if __name__ == '__main__': for tracker_count in range(0, 6): select_window = image.copy() print(f"Add tracker #{tracker_count}") - ct = CornerTracker(do_track, colors[tracker_count], name=f"Tracker-{tracker_count}") + ct = CornerTracker(ct_params, colors[tracker_count], name=f"Tracker-{tracker_count}") ok = ct.init_reference_frame(select_window) if ok: tracker_list.append(ct) @@ -372,6 +376,7 @@ if __name__ == '__main__': for ct in tracker_list: mean_distance = ct.process(image, image_anno) if mean_distance is not None: + scale = ct_params.scale scaled_distance = (scale*mean_distance[0], scale*mean_distance[1]) dist_min[i] = (min(scaled_distance[0], dist_min[i][0]), min(scaled_distance[1], dist_min[i][1])) dist_max[i] = (max(scaled_distance[0], dist_max[i][0]), max(scaled_distance[1], dist_max[i][1]))