aikenml/data_mining
0
1{2 "cells": [3 {4 "cell_type": "code",5 "execution_count": 1,6 "metadata": {},7 "outputs": [],8 "source": [9 "import os\n",10 "import cv2\n",11 "from SegTracker import SegTracker\n",12 "from model_args import aot_args,sam_args,segtracker_args\n",13 "from PIL import Image\n",14 "from aot_tracker import _palette\n",15 "import numpy as np\n",16 "import torch\n",17 "import imageio\n",18 "import matplotlib.pyplot as plt\n",19 "from scipy.ndimage import binary_dilation\n",20 "import gc\n",21 "def save_prediction(pred_mask,output_dir,file_name):\n",22 " save_mask = Image.fromarray(pred_mask.astype(np.uint8))\n",23 " save_mask = save_mask.convert(mode='P')\n",24 " save_mask.putpalette(_palette)\n",25 " save_mask.save(os.path.join(output_dir,file_name))\n",26 "def colorize_mask(pred_mask):\n",27 " save_mask = Image.fromarray(pred_mask.astype(np.uint8))\n",28 " save_mask = save_mask.convert(mode='P')\n",29 " save_mask.putpalette(_palette)\n",30 " save_mask = save_mask.convert(mode='RGB')\n",31 " return np.array(save_mask)\n",32 "def draw_mask(img, mask, alpha=0.5, id_countour=False):\n",33 " img_mask = np.zeros_like(img)\n",34 " img_mask = img\n",35 " if id_countour:\n",36 " # very slow ~ 1s per image\n",37 " obj_ids = np.unique(mask)\n",38 " obj_ids = obj_ids[obj_ids!=0]\n",39 "\n",40 " for id in obj_ids:\n",41 " # Overlay color on binary mask\n",42 " if id <= 255:\n",43 " color = _palette[id*3:id*3+3]\n",44 " else:\n",45 " color = [0,0,0]\n",46 " foreground = img * (1-alpha) + np.ones_like(img) * alpha * np.array(color)\n",47 " binary_mask = (mask == id)\n",48 "\n",49 " # Compose image\n",50 " img_mask[binary_mask] = foreground[binary_mask]\n",51 "\n",52 " countours = binary_dilation(binary_mask,iterations=1) ^ binary_mask\n",53 " img_mask[countours, :] = 0\n",54 " else:\n",55 " binary_mask = (mask!=0)\n",56 " countours = binary_dilation(binary_mask,iterations=1) ^ binary_mask\n",57 " foreground = img*(1-alpha)+colorize_mask(mask)*alpha\n",58 " img_mask[binary_mask] = foreground[binary_mask]\n",59 " img_mask[countours,:] = 0\n",60 " \n",61 " return img_mask.astype(img.dtype)"62 ]63 },64 {65 "attachments": {},66 "cell_type": "markdown",67 "metadata": {},68 "source": [69 "### Set parameters for input and output"70 ]71 },72 {73 "cell_type": "code",74 "execution_count": 2,75 "metadata": {},76 "outputs": [],77 "source": [78 "video_name = 'cell'\n",79 "io_args = {\n",80 " 'input_video': f'./assets/{video_name}.mp4',\n",81 " 'output_mask_dir': f'./assets/{video_name}_masks', # save pred masks\n",82 " 'output_video': f'./assets/{video_name}_seg.mp4', # mask+frame vizualization, mp4 or avi, else the same as input video\n",83 " 'output_gif': f'./assets/{video_name}_seg.gif', # mask visualization\n",84 "}"85 ]86 },87 {88 "attachments": {},89 "cell_type": "markdown",90 "metadata": {},91 "source": [92 "### Tuning SAM on the First Frame for Good Initialization"93 ]94 },95 {96 "cell_type": "code",97 "execution_count": null,98 "metadata": {},99 "outputs": [],100 "source": [101 "# choose good parameters in sam_args based on the first frame segmentation result\n",102 "# other arguments can be modified in model_args.py\n",103 "# note the object number limit is 255 by default, which requires < 10GB GPU memory with amp\n",104 "sam_args['generator_args'] = {\n",105 " 'points_per_side': 30,\n",106 " 'pred_iou_thresh': 0.8,\n",107 " 'stability_score_thresh': 0.9,\n",108 " 'crop_n_layers': 1,\n",109 " 'crop_n_points_downscale_factor': 2,\n",110 " 'min_mask_region_area': 200,\n",111 " }\n",112 "cap = cv2.VideoCapture(io_args['input_video'])\n",113 "frame_idx = 0\n",114 "segtracker = SegTracker(segtracker_args,sam_args,aot_args)\n",115 "segtracker.restart_tracker()\n",116 "with torch.cuda.amp.autocast():\n",117 " while cap.isOpened():\n",118 " ret, frame = cap.read()\n",119 " frame = cv2.cvtColor(frame,cv2.COLOR_BGR2RGB)\n",120 " pred_mask = segtracker.seg(frame)\n",121 " torch.cuda.empty_cache()\n",122 " obj_ids = np.unique(pred_mask)\n",123 " obj_ids = obj_ids[obj_ids!=0]\n",124 " print(\"processed frame {}, obj_num {}\".format(frame_idx,len(obj_ids)),end='\\n')\n",125 " break\n",126 " cap.release()\n",127 " init_res = draw_mask(frame,pred_mask,id_countour=False)\n",128 " plt.figure(figsize=(10,10))\n",129 " plt.axis('off')\n",130 " plt.imshow(init_res)\n",131 " plt.show()\n",132 " plt.figure(figsize=(10,10))\n",133 " plt.axis('off')\n",134 " plt.imshow(colorize_mask(pred_mask))\n",135 " plt.show()\n",136 "\n",137 " del segtracker\n",138 " torch.cuda.empty_cache()\n",139 " gc.collect()"140 ]141 },142 {143 "attachments": {},144 "cell_type": "markdown",145 "metadata": {},146 "source": [147 "### Generate Results for the Whole Video"148 ]149 },150 {151 "cell_type": "code",152 "execution_count": null,153 "metadata": {},154 "outputs": [],155 "source": [156 "# For every sam_gap frames, we use SAM to find new objects and add them for tracking\n",157 "# larger sam_gap is faster but may not spot new objects in time\n",158 "segtracker_args = {\n",159 " 'sam_gap': 5, # the interval to run sam to segment new objects\n",160 " 'min_area': 200, # minimal mask area to add a new mask as a new object\n",161 " 'max_obj_num': 255, # maximal object number to track in a video\n",162 " 'min_new_obj_iou': 0.8, # the area of a new object in the background should > 80% \n",163 "}\n",164 "\n",165 "# source video to segment\n",166 "cap = cv2.VideoCapture(io_args['input_video'])\n",167 "fps = cap.get(cv2.CAP_PROP_FPS)\n",168 "# output masks\n",169 "output_dir = io_args['output_mask_dir']\n",170 "if not os.path.exists(output_dir):\n",171 " os.makedirs(output_dir)\n",172 "pred_list = []\n",173 "masked_pred_list = []\n",174 "\n",175 "torch.cuda.empty_cache()\n",176 "gc.collect()\n",177 "sam_gap = segtracker_args['sam_gap']\n",178 "frame_idx = 0\n",179 "segtracker = SegTracker(segtracker_args,sam_args,aot_args)\n",180 "segtracker.restart_tracker()\n",181 "\n",182 "with torch.cuda.amp.autocast():\n",183 " while cap.isOpened():\n",184 " ret, frame = cap.read()\n",185 " if not ret:\n",186 " break\n",187 " frame = cv2.cvtColor(frame,cv2.COLOR_BGR2RGB)\n",188 " if frame_idx == 0:\n",189 " pred_mask = segtracker.seg(frame)\n",190 " torch.cuda.empty_cache()\n",191 " gc.collect()\n",192 " segtracker.add_reference(frame, pred_mask)\n",193 " elif (frame_idx % sam_gap) == 0:\n",194 " seg_mask = segtracker.seg(frame)\n",195 " torch.cuda.empty_cache()\n",196 " gc.collect()\n",197 " track_mask = segtracker.track(frame)\n",198 " # find new objects, and update tracker with new objects\n",199 " new_obj_mask = segtracker.find_new_objs(track_mask,seg_mask)\n",200 " save_prediction(new_obj_mask,output_dir,str(frame_idx)+'_new.png')\n",201 " pred_mask = track_mask + new_obj_mask\n",202 " # segtracker.restart_tracker()\n",203 " segtracker.add_reference(frame, pred_mask)\n",204 " else:\n",205 " pred_mask = segtracker.track(frame,update_memory=True)\n",206 " torch.cuda.empty_cache()\n",207 " gc.collect()\n",208 " save_prediction(pred_mask,output_dir,str(frame_idx)+'.png')\n",209 " # masked_frame = draw_mask(frame,pred_mask)\n",210 " # masked_pred_list.append(masked_frame)\n",211 " # plt.imshow(masked_frame)\n",212 " # plt.show() \n",213 " \n",214 " pred_list.append(pred_mask)\n",215 " \n",216 " \n",217 " print(\"processed frame {}, obj_num {}\".format(frame_idx,segtracker.get_obj_num()),end='\\r')\n",218 " frame_idx += 1\n",219 " cap.release()\n",220 " print('\\nfinished')"221 ]222 },223 {224 "attachments": {},225 "cell_type": "markdown",226 "metadata": {},227 "source": [228 "### Save results for visualization"229 ]230 },231 {232 "cell_type": "code",233 "execution_count": null,234 "metadata": {},235 "outputs": [],236 "source": [237 "# draw pred mask on frame and save as a video\n",238 "cap = cv2.VideoCapture(io_args['input_video'])\n",239 "fps = cap.get(cv2.CAP_PROP_FPS)\n",240 "width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))\n",241 "height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))\n",242 "num_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))\n",243 "\n",244 "if io_args['input_video'][-3:]=='mp4':\n",245 " fourcc = cv2.VideoWriter_fourcc(*\"mp4v\")\n",246 "elif io_args['input_video'][-3:] == 'avi':\n",247 " fourcc = cv2.VideoWriter_fourcc(*\"MJPG\")\n",248 " # fourcc = cv2.VideoWriter_fourcc(*\"XVID\")\n",249 "else:\n",250 " fourcc = int(cap.get(cv2.CAP_PROP_FOURCC))\n",251 "out = cv2.VideoWriter(io_args['output_video'], fourcc, fps, (width, height))\n",252 "\n",253 "frame_idx = 0\n",254 "while cap.isOpened():\n",255 " ret, frame = cap.read()\n",256 " if not ret:\n",257 " break\n",258 " frame = cv2.cvtColor(frame,cv2.COLOR_BGR2RGB)\n",259 " pred_mask = pred_list[frame_idx]\n",260 " masked_frame = draw_mask(frame,pred_mask)\n",261 " # masked_frame = masked_pred_list[frame_idx]\n",262 " masked_frame = cv2.cvtColor(masked_frame,cv2.COLOR_RGB2BGR)\n",263 " out.write(masked_frame)\n",264 " print('frame {} writed'.format(frame_idx),end='\\r')\n",265 " frame_idx += 1\n",266 "out.release()\n",267 "cap.release()\n",268 "print(\"\\n{} saved\".format(io_args['output_video']))\n",269 "print('\\nfinished')"270 ]271 },272 {273 "cell_type": "code",274 "execution_count": null,275 "metadata": {},276 "outputs": [],277 "source": [278 "# save colorized masks as a gif\n",279 "imageio.mimsave(io_args['output_gif'],pred_list,fps=fps)\n",280 "print(\"{} saved\".format(io_args['output_gif']))"281 ]282 },283 {284 "cell_type": "code",285 "execution_count": 6,286 "metadata": {},287 "outputs": [288 {289 "data": {290 "text/plain": [291 "301"292 ]293 },294 "execution_count": 6,295 "metadata": {},296 "output_type": "execute_result"297 }298 ],299 "source": [300 "# manually release memory (after cuda out of memory)\n",301 "del segtracker\n",302 "torch.cuda.empty_cache()\n",303 "gc.collect()"304 ]305 }306 ],307 "metadata": {308 "kernelspec": {309 "display_name": "Python 3.8.5 64-bit ('ldm': conda)",310 "language": "python",311 "name": "python3"312 },313 "language_info": {314 "codemirror_mode": {315 "name": "ipython",316 "version": 3317 },318 "file_extension": ".py",319 "mimetype": "text/x-python",320 "name": "python",321 "nbconvert_exporter": "python",322 "pygments_lexer": "ipython3",323 "version": "3.8.5"324 },325 "orig_nbformat": 4,326 "vscode": {327 "interpreter": {328 "hash": "536611da043600e50719c9460971b5220bad26cd4a87e5994bfd4c9e9e5e7fb0"329 }330 }331 },332 "nbformat": 4,333 "nbformat_minor": 2334}335 