{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Whale Tail Segment Anything Create Mask","metadata":{"papermill":{"duration":0.005717,"end_time":"2023-05-03T16:39:06.757181","exception":false,"start_time":"2023-05-03T16:39:06.751464","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"https://github.com/facebookresearch/segment-anything<br/>\nhttps://arxiv.org/abs/2304.02643<br/>\nThe Segment Anything (SA) project introduces a new task, model, and dataset for image segmentation. The Segment Anything Model (SAM) and corresponding dataset (SA-1B) are being released to foster research into foundation models for computer vision.","metadata":{"papermill":{"duration":0.003699,"end_time":"2023-05-03T16:39:06.765216","exception":false,"start_time":"2023-05-03T16:39:06.761517","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install git+https://github.com/facebookresearch/segment-anything.git","metadata":{"papermill":{"duration":17.346031,"end_time":"2023-05-03T16:39:24.115135","exception":false,"start_time":"2023-05-03T16:39:06.769104","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-04T04:38:46.283340Z","iopub.execute_input":"2023-05-04T04:38:46.284564Z","iopub.status.idle":"2023-05-04T04:39:03.771919Z","shell.execute_reply.started":"2023-05-04T04:38:46.284512Z","shell.execute_reply":"2023-05-04T04:39:03.769864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\nimport sys\nimport random\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nfrom segment_anything import sam_model_registry \nfrom segment_anything import SamAutomaticMaskGenerator\nfrom segment_anything import SamPredictor","metadata":{"papermill":{"duration":3.49656,"end_time":"2023-05-03T16:39:27.616679","exception":false,"start_time":"2023-05-03T16:39:24.120119","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-04T04:39:03.774338Z","iopub.execute_input":"2023-05-04T04:39:03.775628Z","iopub.status.idle":"2023-05-04T04:39:07.007966Z","shell.execute_reply.started":"2023-05-04T04:39:03.775569Z","shell.execute_reply":"2023-05-04T04:39:07.006593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_anns(anns, axes=None):\n    if len(anns) == 0:\n        return\n    if axes:\n        ax = axes\n    else:\n        ax = plt.gca()\n        ax.set_autoscale_on(False)\n    sorted_anns = sorted(anns, key=(lambda x: x['area']), reverse=True)\n    polygons = []\n    color = []\n    for ann in sorted_anns:\n        m = ann['segmentation']\n        img = np.ones((m.shape[0], m.shape[1], 3))\n        color_mask = np.random.random((1, 3)).tolist()[0]\n        for i in range(3):\n            img[:,:,i] = color_mask[i]\n        ax.imshow(np.dstack((img, m*0.5)))\n\ndef show_mask(mask, ax, random_color=False):\n    if random_color:\n        color = np.concatenate([np.random.random(3), np.array([0.6])], axis=0)\n    else:\n        color = np.array([30/255, 144/255, 255/255, 1.0]) #no transparency\n    h, w = mask.shape[-2:]\n    mask_image = mask.reshape(h, w, 1) * color.reshape(1, 1, -1)\n    ax.imshow(mask_image)\n\n    \ndef show_points(coords, labels, ax, marker_size=375):\n    pos_points = coords[labels==1]\n    neg_points = coords[labels==0]\n    ax.scatter(pos_points[:, 0], pos_points[:, 1], color='green', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)\n    ax.scatter(neg_points[:, 0], neg_points[:, 1], color='red', marker='*', s=marker_size, edgecolor='white', linewidth=1.25)   \n\n    \ndef show_box(box, ax):\n    x0, y0 = box[0], box[1]\n    w, h = box[2] - box[0], box[3] - box[1]\n    ax.add_patch(plt.Rectangle((x0, y0), w, h, edgecolor='green', facecolor=(0,0,0,0), lw=2))    ","metadata":{"papermill":{"duration":0.02585,"end_time":"2023-05-03T16:39:27.647334","exception":false,"start_time":"2023-05-03T16:39:27.621484","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-04T04:39:07.010665Z","iopub.execute_input":"2023-05-04T04:39:07.011328Z","iopub.status.idle":"2023-05-04T04:39:07.029021Z","shell.execute_reply.started":"2023-05-04T04:39:07.011287Z","shell.execute_reply":"2023-05-04T04:39:07.027195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Fill the background with mask other than the tail","metadata":{}},{"cell_type":"code","source":"path0='/kaggle/input/humpback-whale-identification/test/00744bd58.jpg'\nimage = cv2.imread(path0)\nimage = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\nimage = cv2.resize(image,dsize=(200,200))\nplt.figure(figsize=(4,4))\nplt.imshow(image)\nplt.axis('off')\nplt.show()","metadata":{"papermill":{"duration":0.61618,"end_time":"2023-05-03T16:39:28.268245","exception":false,"start_time":"2023-05-03T16:39:27.652065","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-04T04:39:07.031901Z","iopub.execute_input":"2023-05-04T04:39:07.032558Z","iopub.status.idle":"2023-05-04T04:39:07.294739Z","shell.execute_reply.started":"2023-05-04T04:39:07.032517Z","shell.execute_reply":"2023-05-04T04:39:07.293505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# SAM model","metadata":{"papermill":{"duration":0.01122,"end_time":"2023-05-03T16:39:28.29118","exception":false,"start_time":"2023-05-03T16:39:28.27996","status":"completed"},"tags":[]}},{"cell_type":"code","source":"sam_checkpoint = \"/kaggle/input/segment-anything-models/sam_vit_h_4b8939.pth\"\nmodel_type = \"vit_h\"#\ndevice = \"cpu\"\n\nsam = sam_model_registry[model_type](checkpoint=sam_checkpoint)\nsam.to(device=device)\npredictor = SamPredictor(sam)\npredictor.set_image(image)","metadata":{"_kg_hide-output":true,"papermill":{"duration":40.086452,"end_time":"2023-05-03T16:40:08.389225","exception":false,"start_time":"2023-05-03T16:39:28.302773","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-04T04:39:07.296180Z","iopub.execute_input":"2023-05-04T04:39:07.296761Z","iopub.status.idle":"2023-05-04T04:41:12.026762Z","shell.execute_reply.started":"2023-05-04T04:39:07.296715Z","shell.execute_reply":"2023-05-04T04:41:12.024125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# When set label 0 (fish) and label 1 (background)\n### This method is available when the tail position is known","metadata":{}},{"cell_type":"code","source":"print(image.shape)\n\ninput_point = np.array([[5,5],[195,195],[5,195],[195,5],[100,120],[110,110]])\ninput_label = np.array([1,1,1,1,0,0]) # label 1 (green) segmented, label 0 (red) excluded\n\nplt.figure(figsize=(4,4))\nplt.imshow(image)\nshow_points(input_point, input_label, plt.gca())\nplt.axis('on')\nplt.title('Original', fontsize=12)\nplt.show()  \n\nmasks, scores, logits = predictor.predict(\n    point_coords=input_point,\n    point_labels=input_label,\n    multimask_output=True,\n)\n\nfor i, (mask, score) in enumerate(zip(masks, scores)):\n    plt.figure(figsize=(4,4))\n    plt.imshow(image)\n    show_mask(mask, plt.gca())\n    #show_points(input_point, input_label, plt.gca())\n    plt.title(f\"Mask {i+1}, Score: {score:.3f}\", fontsize=12)\n    plt.show()  \n  ","metadata":{"papermill":{"duration":83.996051,"end_time":"2023-05-03T16:41:32.457309","exception":false,"start_time":"2023-05-03T16:40:08.461258","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-04T04:42:50.006337Z","iopub.execute_input":"2023-05-04T04:42:50.006814Z","iopub.status.idle":"2023-05-04T04:42:51.488752Z","shell.execute_reply.started":"2023-05-04T04:42:50.006774Z","shell.execute_reply":"2023-05-04T04:42:51.487519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# When set only label 0 (background)\n### This method is availbale even when the tail position is unknwon ","metadata":{}},{"cell_type":"code","source":"input_point2 = np.array([[5,5],[195,195],[5,195],[195,5]])\ninput_label2 = np.array([1,1,1,1]) # label 1 (green) segmented, label 0 (red) excluded\n\nplt.figure(figsize=(4,4))\nplt.imshow(image)\nshow_points(input_point2, input_label2, plt.gca())\nplt.axis('on')\nplt.title('Original', fontsize=12)\nplt.show()  \n\nmasks2, scores2, logits2 = predictor.predict(\n    point_coords=input_point2,\n    point_labels=input_label2,\n    multimask_output=True,\n)\n\n    \nfor i, (mask, score) in enumerate(zip(masks2, scores2)):\n    plt.figure(figsize=(4,4))\n    plt.imshow(image)\n    show_mask(mask, plt.gca())\n    #show_points(input_point, input_label, plt.gca())\n    plt.title(f\"Mask {i+1}, Score: {score:.3f}\", fontsize=12)\n    plt.savefig(f\"mask{i+1}.png\")\n    plt.show()    ","metadata":{"papermill":{"duration":0.039804,"end_time":"2023-05-03T16:45:29.508336","exception":false,"start_time":"2023-05-03T16:45:29.468532","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-04T04:41:13.227000Z","iopub.execute_input":"2023-05-04T04:41:13.227726Z","iopub.status.idle":"2023-05-04T04:41:14.394717Z","shell.execute_reply.started":"2023-05-04T04:41:13.227687Z","shell.execute_reply":"2023-05-04T04:41:14.393324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}