CoolFace
Apppublic

aikenml/data_mining

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
SegTracker.py264 linesDownload Raw Back to root
1import sys2sys.path.append("..")3sys.path.append("./sam")4from sam.segment_anything import sam_model_registry, SamAutomaticMaskGenerator5from aot_tracker import get_aot6import numpy as np7from tool.segmentor import Segmentor8from tool.detector import Detector9from tool.transfer_tools import draw_outline, draw_points10import cv211from seg_track_anything import draw_mask12 13 14class SegTracker():15    def __init__(self,segtracker_args, sam_args, aot_args) -> None:16        """17         Initialize SAM and AOT.18        """19        self.sam = Segmentor(sam_args)20        self.tracker = get_aot(aot_args)21        self.detector = Detector(self.sam.device)22        self.sam_gap = segtracker_args['sam_gap']23        self.min_area = segtracker_args['min_area']24        self.max_obj_num = segtracker_args['max_obj_num']25        self.min_new_obj_iou = segtracker_args['min_new_obj_iou']26        self.reference_objs_list = []27        self.object_idx = 128        self.curr_idx = 129        self.origin_merged_mask = None  # init by segment-everything or update30        self.first_frame_mask = None31 32        # debug33        self.everything_points = []34        self.everything_labels = []35        print("SegTracker has been initialized")36 37    def seg(self,frame):38        '''39        Arguments:40            frame: numpy array (h,w,3)41        Return:42            origin_merged_mask: numpy array (h,w)43        '''44        frame = frame[:, :, ::-1]45        anns = self.sam.everything_generator.generate(frame)46 47        # anns is a list recording all predictions in an image48        if len(anns) == 0:49            return50        # merge all predictions into one mask (h,w)51        # note that the merged mask may lost some objects due to the overlapping52        self.origin_merged_mask = np.zeros(anns[0]['segmentation'].shape,dtype=np.uint8)53        idx = 154        for ann in anns:55            if ann['area'] > self.min_area:56                m = ann['segmentation']57                self.origin_merged_mask[m==1] = idx58                idx += 159                self.everything_points.append(ann["point_coords"][0])60                self.everything_labels.append(1)61 62        obj_ids = np.unique(self.origin_merged_mask)63        obj_ids = obj_ids[obj_ids!=0]64 65        self.object_idx = 166        for id in obj_ids:67            if np.sum(self.origin_merged_mask==id) < self.min_area or self.object_idx > self.max_obj_num:68                self.origin_merged_mask[self.origin_merged_mask==id] = 069            else:70                self.origin_merged_mask[self.origin_merged_mask==id] = self.object_idx71                self.object_idx += 172 73        self.first_frame_mask = self.origin_merged_mask74        return self.origin_merged_mask75 76    def update_origin_merged_mask(self, updated_merged_mask):77        self.origin_merged_mask = updated_merged_mask78        # obj_ids = np.unique(updated_merged_mask)79        # obj_ids = obj_ids[obj_ids!=0]80        # self.object_idx = int(max(obj_ids)) + 181 82    def reset_origin_merged_mask(self, mask, id):83        self.origin_merged_mask = mask84        self.curr_idx = id85 86    def add_reference(self,frame,mask,frame_step=0):87        '''88        Add objects in a mask for tracking.89        Arguments:90            frame: numpy array (h,w,3)91            mask: numpy array (h,w)92        '''93        self.reference_objs_list.append(np.unique(mask))94        self.curr_idx = self.get_obj_num() + 195        self.tracker.add_reference_frame(frame,mask, self.curr_idx - 1, frame_step)96 97    def track(self,frame,update_memory=False):98        '''99        Track all known objects.100        Arguments:101            frame: numpy array (h,w,3)102        Return:103            origin_merged_mask: numpy array (h,w)104        '''105        pred_mask = self.tracker.track(frame)106        if update_memory:107            self.tracker.update_memory(pred_mask)108        return pred_mask.squeeze(0).squeeze(0).detach().cpu().numpy().astype(np.uint8)109    110    def get_tracking_objs(self):111        objs = set()112        for ref in self.reference_objs_list:113            objs.update(set(ref))114        objs = list(sorted(list(objs)))115        objs = [i for i in objs if i!=0]116        return objs117    118    def get_obj_num(self):119        objs = self.get_tracking_objs()120        if len(objs) == 0: return 0121        return int(max(objs))122 123    def find_new_objs(self, track_mask, seg_mask):124        '''125        Compare tracked results from AOT with segmented results from SAM. Select objects from background if they are not tracked.126        Arguments:127            track_mask: numpy array (h,w)128            seg_mask: numpy array (h,w)129        Return:130            new_obj_mask: numpy array (h,w)131        '''132        new_obj_mask = (track_mask==0) * seg_mask133        new_obj_ids = np.unique(new_obj_mask)134        new_obj_ids = new_obj_ids[new_obj_ids!=0]135        # obj_num = self.get_obj_num() + 1136        obj_num = self.curr_idx137        for idx in new_obj_ids:138            new_obj_area = np.sum(new_obj_mask==idx)139            obj_area = np.sum(seg_mask==idx)140            if new_obj_area/obj_area < self.min_new_obj_iou or new_obj_area < self.min_area\141                or obj_num > self.max_obj_num:142                new_obj_mask[new_obj_mask==idx] = 0143            else:144                new_obj_mask[new_obj_mask==idx] = obj_num145                obj_num += 1146        return new_obj_mask147        148    def restart_tracker(self):149        self.tracker.restart()150 151    def seg_acc_bbox(self, origin_frame: np.ndarray, bbox: np.ndarray,):152        ''''153        Use bbox-prompt to get mask154        Parameters:155            origin_frame: H, W, C156            bbox: [[x0, y0], [x1, y1]]157        Return:158            refined_merged_mask: numpy array (h, w)159            masked_frame: numpy array (h, w, c)160        '''161        # get interactive_mask162        interactive_mask = self.sam.segment_with_box(origin_frame, bbox)[0]163        refined_merged_mask = self.add_mask(interactive_mask)164 165        # draw mask166        masked_frame = draw_mask(origin_frame.copy(), refined_merged_mask)167 168        # draw bbox169        masked_frame = cv2.rectangle(masked_frame, bbox[0], bbox[1], (0, 0, 255))170 171        return refined_merged_mask, masked_frame172 173    def seg_acc_click(self, origin_frame: np.ndarray, coords: np.ndarray, modes: np.ndarray, multimask=True):174        '''175        Use point-prompt to get mask176        Parameters:177            origin_frame: H, W, C178            coords: nd.array [[x, y]]179            modes: nd.array [[1]]180        Return:181            refined_merged_mask: numpy array (h, w)182            masked_frame: numpy array (h, w, c)183        '''184        # get interactive_mask185        interactive_mask = self.sam.segment_with_click(origin_frame, coords, modes, multimask)186 187        refined_merged_mask = self.add_mask(interactive_mask)188 189        # draw mask190        masked_frame = draw_mask(origin_frame.copy(), refined_merged_mask)191 192        # draw points193        # self.everything_labels = np.array(self.everything_labels).astype(np.int64)194        # self.everything_points = np.array(self.everything_points).astype(np.int64)195 196        masked_frame = draw_points(coords, modes, masked_frame)197 198        # draw outline199        masked_frame = draw_outline(interactive_mask, masked_frame)200 201        return refined_merged_mask, masked_frame202 203    def add_mask(self, interactive_mask: np.ndarray):204        '''205        Merge interactive mask with self.origin_merged_mask206        Parameters:207            interactive_mask: numpy array (h, w)208        Return:209            refined_merged_mask: numpy array (h, w)210        '''211        if self.origin_merged_mask is None:212            self.origin_merged_mask = np.zeros(interactive_mask.shape,dtype=np.uint8)213 214        refined_merged_mask = self.origin_merged_mask.copy()215        refined_merged_mask[interactive_mask > 0] = self.curr_idx216 217        return refined_merged_mask218    219    def detect_and_seg(self, origin_frame: np.ndarray, grounding_caption, box_threshold, text_threshold, box_size_threshold=1, reset_image=False):220        '''221        Using Grounding-DINO to detect object acc Text-prompts222        Retrun:223            refined_merged_mask: numpy array (h, w)224            annotated_frame: numpy array (h, w, 3)225        '''226        # backup id and origin-merged-mask227        bc_id = self.curr_idx228        bc_mask = self.origin_merged_mask229 230        # get annotated_frame and boxes231        annotated_frame, boxes = self.detector.run_grounding(origin_frame, grounding_caption, box_threshold, text_threshold)232        for i in range(len(boxes)):233            bbox = boxes[i]234            if (bbox[1][0] - bbox[0][0]) * (bbox[1][1] - bbox[0][1]) > annotated_frame.shape[0] * annotated_frame.shape[1] * box_size_threshold:235                continue236            interactive_mask = self.sam.segment_with_box(origin_frame, bbox, reset_image)[0]237            refined_merged_mask = self.add_mask(interactive_mask)238            self.update_origin_merged_mask(refined_merged_mask)239            self.curr_idx += 1240 241        # reset origin_mask242        self.reset_origin_merged_mask(bc_mask, bc_id)243 244        return refined_merged_mask, annotated_frame245 246if __name__ == '__main__':247    from model_args import segtracker_args,sam_args,aot_args248 249    Seg_Tracker = SegTracker(segtracker_args, sam_args, aot_args)250    251    # ------------------ detect test ----------------------252    253    origin_frame = cv2.imread('/data2/cym/Seg_Tra_any/Segment-and-Track-Anything/debug/point.png')254    origin_frame = cv2.cvtColor(origin_frame, cv2.COLOR_BGR2RGB)255    grounding_caption = "swan.water"256    box_threshold = 0.25257    text_threshold = 0.25258 259    predicted_mask, annotated_frame = Seg_Tracker.detect_and_seg(origin_frame, grounding_caption, box_threshold, text_threshold)260    masked_frame = draw_mask(annotated_frame, predicted_mask)261    origin_frame = cv2.cvtColor(origin_frame, cv2.COLOR_RGB2BGR)262 263    cv2.imwrite('./debug/masked_frame.png', masked_frame)264    cv2.imwrite('./debug/x.png', annotated_frame)