CoolFace
Apppublic

aikenml/data_mining

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
demo.ipynb335 linesDownload Raw Back to root
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