{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Segment Anything Model (SAM)\nResearch by Meta AI\n\nhttps://segment-anything.com/\n> SAM is a promptable segmentation system with zero-shot generalization to unfamiliar objects and images, without the need for additional training.\n\nhttps://github.com/facebookresearch/segment-anything\n\n> The Segment Anything Model (SAM) produces high quality object masks from input prompts such as points or boxes, and it can be used to generate masks for all objects in an image. It has been trained on a dataset of 11 million images and 1.1 billion masks, and has strong zero-shot performance on a variety of segmentation tasks.","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nfrom matplotlib import pyplot as plt\nimport torch\nimport cv2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-07T07:17:41.41619Z","iopub.execute_input":"2023-09-07T07:17:41.416524Z","iopub.status.idle":"2023-09-07T07:17:44.874902Z","shell.execute_reply.started":"2023-09-07T07:17:41.416491Z","shell.execute_reply":"2023-09-07T07:17:44.873757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install git+https://github.com/facebookresearch/segment-anything.git","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:17:44.877223Z","iopub.execute_input":"2023-09-07T07:17:44.878058Z","iopub.status.idle":"2023-09-07T07:18:02.079282Z","shell.execute_reply.started":"2023-09-07T07:17:44.878011Z","shell.execute_reply":"2023-09-07T07:18:02.077961Z"},"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, 0.6])\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":{"execution":{"iopub.status.busy":"2023-09-07T07:18:02.080864Z","iopub.execute_input":"2023-09-07T07:18:02.081248Z","iopub.status.idle":"2023-09-07T07:18:02.102687Z","shell.execute_reply.started":"2023-09-07T07:18:02.081212Z","shell.execute_reply":"2023-09-07T07:18:02.099992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Automatically Generating Object Masks with SAM\n\nhttps://github.com/facebookresearch/segment-anything/blob/main/notebooks/automatic_mask_generator_example.ipynb","metadata":{}},{"cell_type":"code","source":"from segment_anything import sam_model_registry, SamAutomaticMaskGenerator, SamPredictor\n\nsam_checkpoint = \"/kaggle/input/segment-anything/pytorch/vit-b/1/model.pth\"\nmodel_type = \"vit_b\"\n\ndevice = \"cuda\"\n\nsam = sam_model_registry[model_type](checkpoint=sam_checkpoint)\nsam.to(device=device)\n\nmask_generator = SamAutomaticMaskGenerator(sam, points_per_batch=16)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:18:02.106039Z","iopub.execute_input":"2023-09-07T07:18:02.106694Z","iopub.status.idle":"2023-09-07T07:18:12.894046Z","shell.execute_reply.started":"2023-09-07T07:18:02.106655Z","shell.execute_reply":"2023-09-07T07:18:12.892904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '/kaggle/input/stable-diffusion-image-to-prompts/images/c98f79f71.png'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\nmasks = mask_generator.generate(image_array)\n\n_, axes = plt.subplots(1,3, figsize=(16,16))\naxes[0].imshow(image_array)\nshow_anns(masks, axes[1])\naxes[2].imshow(image_array)\nshow_anns(masks, axes[2])","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:18:12.895581Z","iopub.execute_input":"2023-09-07T07:18:12.895962Z","iopub.status.idle":"2023-09-07T07:18:23.473204Z","shell.execute_reply.started":"2023-09-07T07:18:12.895905Z","shell.execute_reply":"2023-09-07T07:18:23.471876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '/kaggle/input/the-car-connection-picture-dataset/Acura_ILX_2014_28_16_110_15_4_70_55_179_39_FWD_5_4_4dr_Klc.jpg'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\nmasks = mask_generator.generate(image_array)\n\n_, axes = plt.subplots(1,3, figsize=(16,16))\naxes[0].imshow(image_array)\nshow_anns(masks, axes[1])\naxes[2].imshow(image_array)\nshow_anns(masks, axes[2])","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:18:23.474284Z","iopub.execute_input":"2023-09-07T07:18:23.474635Z","iopub.status.idle":"2023-09-07T07:18:31.277227Z","shell.execute_reply.started":"2023-09-07T07:18:23.474601Z","shell.execute_reply":"2023-09-07T07:18:31.275953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '/kaggle/input/happy-whale-and-dolphin/train_images/00177f3c614d1e.jpg'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\nmasks = mask_generator.generate(image_array)\n\n_, axes = plt.subplots(1,3, figsize=(16,16))\naxes[0].imshow(image_array)\nshow_anns(masks, axes[1])\naxes[2].imshow(image_array)\nshow_anns(masks, axes[2])","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:18:31.278972Z","iopub.execute_input":"2023-09-07T07:18:31.27971Z","iopub.status.idle":"2023-09-07T07:19:27.385737Z","shell.execute_reply.started":"2023-09-07T07:18:31.279663Z","shell.execute_reply":"2023-09-07T07:19:27.384767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '/kaggle/input/food-segmentation/Food Segmentation/images/8.png'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\nmasks = mask_generator.generate(image_array)\n\n_, axes = plt.subplots(1,3, figsize=(16,16))\naxes[0].imshow(image_array)\nshow_anns(masks, axes[1])\naxes[2].imshow(image_array)\nshow_anns(masks, axes[2])","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:19:27.387219Z","iopub.execute_input":"2023-09-07T07:19:27.38835Z","iopub.status.idle":"2023-09-07T07:20:05.168697Z","shell.execute_reply.started":"2023-09-07T07:19:27.388305Z","shell.execute_reply":"2023-09-07T07:20:05.167735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del(mask_generator)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:20:05.170304Z","iopub.execute_input":"2023-09-07T07:20:05.170984Z","iopub.status.idle":"2023-09-07T07:20:05.176202Z","shell.execute_reply.started":"2023-09-07T07:20:05.17094Z","shell.execute_reply":"2023-09-07T07:20:05.174996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Object masks from prompts with SAM\n\nhttps://github.com/facebookresearch/segment-anything/blob/main/notebooks/predictor_example.ipynb","metadata":{}},{"cell_type":"code","source":"from segment_anything import sam_model_registry, SamPredictor\n\npredictor = SamPredictor(sam)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:20:05.180368Z","iopub.execute_input":"2023-09-07T07:20:05.180816Z","iopub.status.idle":"2023-09-07T07:20:05.193024Z","shell.execute_reply.started":"2023-09-07T07:20:05.180779Z","shell.execute_reply":"2023-09-07T07:20:05.191398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Selecting objects with SAM\n","metadata":{}},{"cell_type":"markdown","source":"## Specifying a specific object with points","metadata":{}},{"cell_type":"code","source":"image_path = '/kaggle/input/the-car-connection-picture-dataset/Acura_ILX_2014_28_16_110_15_4_70_55_179_39_FWD_5_4_4dr_Klc.jpg'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\npredictor.set_image(image_array)\n\ninput_point = np.array([[120, 135]])\ninput_label = np.array([1])\n\nplt.imshow(image_array)\nshow_points(input_point, input_label, plt.gca())\nplt.axis('on')\nplt.show()  \n\n\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=(10,10))\n    plt.imshow(image_array)\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=18)\n    plt.show()  \n  ","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:20:05.194818Z","iopub.execute_input":"2023-09-07T07:20:05.195662Z","iopub.status.idle":"2023-09-07T07:20:06.851975Z","shell.execute_reply.started":"2023-09-07T07:20:05.195624Z","shell.execute_reply":"2023-09-07T07:20:06.850935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Specifying a specific object with additional points (include and exclude)","metadata":{}},{"cell_type":"code","source":"image_path = '/kaggle/input/stable-diffusion-image-to-prompts/images/c98f79f71.png'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\npredictor.set_image(image_array)\n\ninput_point = np.array([[300, 50], [185, 100], [250, 200], [400, 50]])\ninput_label = np.array([1, 1, 0, 1])\n\nplt.imshow(image_array)\nshow_points(input_point, input_label, plt.gca())\nplt.axis('on')\nplt.show()  \n\n\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=(10,10))\n    plt.imshow(image_array)\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=18)\n    plt.show()  ","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:20:06.853762Z","iopub.execute_input":"2023-09-07T07:20:06.854468Z","iopub.status.idle":"2023-09-07T07:20:08.568989Z","shell.execute_reply.started":"2023-09-07T07:20:06.854419Z","shell.execute_reply":"2023-09-07T07:20:08.567913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '/kaggle/input/happy-whale-and-dolphin/train_images/00177f3c614d1e.jpg'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\npredictor.set_image(image_array)\n\ninput_point = np.array([[1400, 1600], [2500, 500]])\ninput_label = np.array([1, 0])\n\nmask_input = logits[np.argmax(scores), :, :]  # Choose the model's best mask\nmasks, _, _ = predictor.predict(\n    point_coords=input_point,\n    point_labels=input_label,\n    mask_input=mask_input[None, :, :],\n    multimask_output=False,\n)\n\nplt.imshow(image_array)\nshow_mask(masks, plt.gca())\nshow_points(input_point, input_label, plt.gca())\nplt.show()   ","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:20:08.570572Z","iopub.execute_input":"2023-09-07T07:20:08.571382Z","iopub.status.idle":"2023-09-07T07:20:12.228959Z","shell.execute_reply.started":"2023-09-07T07:20:08.571333Z","shell.execute_reply":"2023-09-07T07:20:12.227759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '/kaggle/input/tongue1/Tongue/Covid-19 tongue/Covid-19_tongue.jpg'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\npredictor.set_image(image_array)\n\ninput_point = np.array([[1400, 1600], [2500, 500]])\ninput_label = np.array([1, 0])\n\nmask_input = logits[np.argmax(scores), :, :]  # Choose the model's best mask\nmasks, _, _ = predictor.predict(\n    point_coords=input_point,\n    point_labels=input_label,\n    mask_input=mask_input[None, :, :],\n    multimask_output=False,\n)\n\nplt.imshow(image_array)\nshow_mask(masks, plt.gca())\nshow_points(input_point, input_label, plt.gca())\nplt.show()  ","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:21:52.537548Z","iopub.execute_input":"2023-09-07T07:21:52.537982Z","iopub.status.idle":"2023-09-07T07:21:53.033013Z","shell.execute_reply.started":"2023-09-07T07:21:52.537934Z","shell.execute_reply":"2023-09-07T07:21:53.031897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_path = '/kaggle/input/food-segmentation/Food Segmentation/images/10.png'\nimage_array = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\npredictor.set_image(image_array)\n\ninput_point = np.array([[1400, 1600], [2500, 500]])\ninput_label = np.array([1, 0])\n\nmask_input = logits[np.argmax(scores), :, :]  # Choose the model's best mask\nmasks, _, _ = predictor.predict(\n    point_coords=input_point,\n    point_labels=input_label,\n    mask_input=mask_input[None, :, :],\n    multimask_output=False,\n)\n\nplt.imshow(image_array)\nshow_mask(masks, plt.gca())\nshow_points(input_point, input_label, plt.gca())\nplt.show() ","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:24:57.321944Z","iopub.execute_input":"2023-09-07T07:24:57.322694Z","iopub.status.idle":"2023-09-07T07:24:58.089962Z","shell.execute_reply.started":"2023-09-07T07:24:57.32265Z","shell.execute_reply":"2023-09-07T07:24:58.088859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Specifying a specific object with a box\n","metadata":{}},{"cell_type":"code","source":"input_box = np.array([600, 1100, 2100, 1800])\nmasks, _, _ = predictor.predict(\n    point_coords=None,\n    point_labels=None,\n    box=input_box[None, :],\n    multimask_output=False,\n)\nplt.imshow(image_array)\nshow_mask(masks[0], plt.gca())\nshow_box(input_box, plt.gca())\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-07T07:20:12.230778Z","iopub.execute_input":"2023-09-07T07:20:12.231184Z","iopub.status.idle":"2023-09-07T07:20:14.952439Z","shell.execute_reply.started":"2023-09-07T07:20:12.231146Z","shell.execute_reply":"2023-09-07T07:20:14.951283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Follow this notebook for additional prompting options:\n\nhttps://github.com/facebookresearch/segment-anything/blob/main/notebooks/predictor_example.ipynb","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"}}