aikenml/data_mining
0
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)