{"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":"# 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-07-25T14:16:02.116091Z","iopub.execute_input":"2023-07-25T14:16:02.116496Z","iopub.status.idle":"2023-07-25T14:16:02.12244Z","shell.execute_reply.started":"2023-07-25T14:16:02.116458Z","shell.execute_reply":"2023-07-25T14:16:02.121134Z"},"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-07-25T14:16:02.151316Z","iopub.execute_input":"2023-07-25T14:16:02.151819Z","iopub.status.idle":"2023-07-25T14:16:17.040859Z","shell.execute_reply.started":"2023-07-25T14:16:02.151785Z","shell.execute_reply":"2023-07-25T14:16:17.039372Z"},"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-07-25T14:16:17.044231Z","iopub.execute_input":"2023-07-25T14:16:17.044693Z","iopub.status.idle":"2023-07-25T14:16:17.063803Z","shell.execute_reply.started":"2023-07-25T14:16:17.044641Z","shell.execute_reply":"2023-07-25T14:16:17.062472Z"},"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-07-25T14:16:17.065714Z","iopub.execute_input":"2023-07-25T14:16:17.066628Z","iopub.status.idle":"2023-07-25T14:16:18.273926Z","shell.execute_reply.started":"2023-07-25T14:16:17.066563Z","shell.execute_reply":"2023-07-25T14:16:18.27263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nkaggle datasets download -d adityaraj01/cityscape-pic-01\nimage_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-07-25T14:36:34.381684Z","iopub.execute_input":"2023-07-25T14:36:34.382205Z","iopub.status.idle":"2023-07-25T14:36:34.420125Z","shell.execute_reply.started":"2023-07-25T14:36:34.382121Z","shell.execute_reply":"2023-07-25T14:36:34.41865Z"},"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-07-25T14:16:24.951985Z","iopub.execute_input":"2023-07-25T14:16:24.952688Z","iopub.status.idle":"2023-07-25T14:16:32.931734Z","shell.execute_reply.started":"2023-07-25T14:16:24.952645Z","shell.execute_reply":"2023-07-25T14:16:32.930656Z"},"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-07-25T14:16:32.933498Z","iopub.execute_input":"2023-07-25T14:16:32.934158Z","iopub.status.idle":"2023-07-25T14:17:30.596148Z","shell.execute_reply.started":"2023-07-25T14:16:32.934118Z","shell.execute_reply":"2023-07-25T14:17:30.594925Z"},"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-07-25T14:17:30.597941Z","iopub.execute_input":"2023-07-25T14:17:30.59873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del(mask_generator)","metadata":{"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":{"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":{"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":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"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":{}}]}