diff --git a/ocv_corner_tracker.py b/ocv_corner_tracker.py index 7c9d406..06ab4f4 100644 --- a/ocv_corner_tracker.py +++ b/ocv_corner_tracker.py @@ -1,6 +1,7 @@ import cv2 import numpy as np import argparse +import os import sys import csv import json @@ -100,8 +101,8 @@ class CornerTrackerSettings: self.tracking_bb = _tracking_bb self.matching_tpl_bb = [] - def from_dict(self, settings: dict): - self.__dict__.update(settings) + def from_dict(self, _settings: dict): + self.__dict__.update(_settings) return self def append(self, _bb: np.array): @@ -164,12 +165,11 @@ class CornerTracker: ''' return _result - def create_from_settings(self, _image: np.array, settings: CornerTrackerSettings): - _mask = CornerTracker.create_mask(_image, settings.mask_bb) - image_processed_local, _ = self.create_tracking(_image, _mask, settings.tracking_bb) - self.create_matching(image_processed_local, settings.tracking_bb, settings.matching_tpl_bb) + def create_from_settings(self, _image: np.array, _settings: CornerTrackerSettings, _image_anno: np.array): + _mask = CornerTracker.create_mask(_image, _settings.mask_bb) + image_processed_local, _ = self.create_tracking(_image, _mask, _settings.tracking_bb) - return self.create_matching(image_processed_local, settings.tracking_bb, settings.matching_tpl_bb) + return self.create_matching(image_processed_local, _settings.tracking_bb, _settings.matching_tpl_bb, _image_anno) @staticmethod def create_mask(_image: np.array, mask_bb: np.array): @@ -199,7 +199,7 @@ class CornerTracker: return image_processed_local, tracking_img - def create_matching(self, image_processed_local: np.array, tracking_bb: np.array, matching_tpl_bb_list: list): + def create_matching(self, image_processed_local: np.array, tracking_bb: np.array, matching_tpl_bb_list: list, _image_anno: np.array): # Create corner matcher self.corner_matcher_list = [] corner_list = [] @@ -226,31 +226,34 @@ class CornerTracker: corner = (corner_local[0] + tracking_bb[0], corner_local[1] + tracking_bb[1]) self.corner_ref.append(corner) + for matching_tpl_bb in matching_tpl_bb_list: + self._debug(_image_anno, matching_tpl_bb) + return True - def init_reference_frame(self, _image: np.array): - settings = CornerTrackerSettings(self.id) + def init_reference_frame(self, _image: np.array, _image_anno: np.array): + _settings = CornerTrackerSettings(self.id) # Draw mask _mask = CornerTracker.mask_init(_image) # Init mask print(f"Draw mask object") - _mask_bb = cv2.selectROI("Tracker Reference", _image, False) + _mask_bb = cv2.selectROI("Tracker Reference", _image_anno, False) cv2.destroyWindow("Tracker Reference") if bbox_center(_mask_bb) == (0, 0): _mask_bb = (0, 0, _image.shape[1], _image.shape[0]) - settings.mask_bb = _mask_bb + _settings.mask_bb = _mask_bb _mask = self.create_mask(_image, _mask_bb) # Select search area print(f"Select tracking object") - tracking_bb = cv2.selectROI("Tracker Reference", _image, False) + tracking_bb = cv2.selectROI("Tracker Reference", _image_anno, False) cv2.destroyWindow("Tracker Reference") if bbox_center(tracking_bb) == (0, 0): return None - settings.tracking_bb = tracking_bb + _settings.tracking_bb = tracking_bb image_processed_local, tracking_img = self.create_tracking(_image, _mask, tracking_bb) count = 1 @@ -260,15 +263,15 @@ class CornerTracker: if bbox_center(matching_tpl_bb) == (0, 0): break - settings.matching_tpl_bb.append(matching_tpl_bb) + _settings.matching_tpl_bb.append(matching_tpl_bb) print(f"Corner {count} added") count += 1 cv2.destroyWindow("Matcher Reference") print(f"Added {count} corners") - if self.create_matching(image_processed_local, settings.tracking_bb, settings.matching_tpl_bb): - return settings + if self.create_matching(image_processed_local, _settings.tracking_bb, _settings.matching_tpl_bb, _image_anno): + return _settings def process(self, _image: np.array, _image_anno: np.array): # Update tracker bb @@ -363,17 +366,17 @@ class CornerTracker: def project_load(path: str = "./"): _settings = None try: - with open(path + "ocv_tracker_prj.json", "r") as fp: + with open(os.path.join(path, "ocv_tracker_prj.json"), "r") as fp: _settings = json.load(fp) - except Exception: + except FileNotFoundError: pass return _settings -def project_save(settings: dict, path: str = "./"): - with open(path + "ocv_tracker_prj.json", "w") as fp: - json.dump(settings, fp, indent=4) +def project_save(_settings: dict, path: str = "./"): + with open(os.path.join(path, "ocv_tracker_prj.json"), "w") as fp: + json.dump(_settings, fp, indent=4) if __name__ == '__main__': @@ -393,7 +396,7 @@ if __name__ == '__main__': # Fill CornerTrackerParams from args ct_params = CornerTrackerParams(_scale=args.scale, _pre_track=args.track, _show_path=args.path) - prj = project_load() + prj = project_load(os.path.dirname(args.filename)) prj_loaded = False if prj is not None: ct_params = ct_params.from_dict(prj['params']) @@ -412,26 +415,27 @@ if __name__ == '__main__': colors = [(255, 0, 0), (0, 255, 0), (0, 0, 255), (255, 0, 255), (0, 255, 255), (255, 255, 255)] tracker_count = 0 tracker_list = [] - select_window = image.copy() + image_anno = image.copy() if prj_loaded: ct_params = CornerTrackerParams().from_dict(prj['params']) for tracker_settings in prj['trackers']: print(tracker_settings) ct = CornerTracker(tracker_settings['id'], ct_params, colors[tracker_count], name=f"Tracker-{tracker_settings['id']}") - ct.create_from_settings(select_window, CornerTrackerSettings(tracker_count).from_dict(tracker_settings)) + ct.create_from_settings(image, CornerTrackerSettings(tracker_count).from_dict(tracker_settings), image_anno) tracker_list.append(ct) tracker_count += 1 else: prj = {'params': ct_params.__dict__, 'trackers': []} + cv2.imshow("Tracker Reference", image_anno) + print("Press 'a' for adding trackers, press 'd' to delete a tracker, press other key to continue") - k = cv2.waitKey(-1) & 0xff while True: print(f"Add tracker #{tracker_count}") ct = CornerTracker(tracker_count, ct_params, colors[tracker_count], name=f"Tracker-{tracker_count}") - settings = ct.init_reference_frame(select_window) + settings = ct.init_reference_frame(image, image_anno) if settings is not None: prj['trackers'].append(settings.__dict__) tracker_list.append(ct) @@ -440,7 +444,7 @@ if __name__ == '__main__': break tracker_count += 1 - project_save(prj) + project_save(prj, os.path.dirname(args.filename)) # rewind to frame 0 video.set(cv2.CAP_PROP_POS_FRAMES, 0) @@ -467,7 +471,6 @@ if __name__ == '__main__': cv2.rectangle(image_anno, (25, 0), (int(image_anno.shape[1]), 25*(1+len(tracker_list))), (0, 0, 0), -1) for ct in tracker_list: mean_distance = ct.process(image, image_anno) - ct.id if mean_distance is not None: scale = np.float32(ct_params.scale) scaled_distance = (scale*mean_distance[0], scale*mean_distance[1])