{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":22962,"databundleVersionId":3171193,"sourceType":"competition"},{"sourceId":45917,"databundleVersionId":5024308,"sourceType":"competition"},{"sourceId":928083,"sourceType":"datasetVersion","datasetId":501015},{"sourceId":6230584,"sourceType":"datasetVersion","datasetId":2982879},{"sourceId":3848,"sourceType":"modelInstanceVersion","modelInstanceId":2749}],"dockerImageVersionId":30461,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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-05-14T13:56:36.165923Z","iopub.execute_input":"2023-05-14T13:56:36.166323Z","iopub.status.idle":"2023-05-14T13:56:38.412461Z","shell.execute_reply.started":"2023-05-14T13:56:36.16629Z","shell.execute_reply":"2023-05-14T13:56:38.411139Z"},"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-05-14T13:56:38.414779Z","iopub.execute_input":"2023-05-14T13:56:38.415349Z","iopub.status.idle":"2023-05-14T13:56:56.168382Z","shell.execute_reply.started":"2023-05-14T13:56:38.415304Z","shell.execute_reply":"2023-05-14T13:56:56.166991Z"},"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-05-14T13:56:56.17109Z","iopub.execute_input":"2023-05-14T13:56:56.171908Z","iopub.status.idle":"2023-05-14T13:56:56.190922Z","shell.execute_reply.started":"2023-05-14T13:56:56.171857Z","shell.execute_reply":"2023-05-14T13:56:56.189253Z"},"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-05-14T13:56:56.19565Z","iopub.execute_input":"2023-05-14T13:56:56.196129Z","iopub.status.idle":"2023-05-14T13:57:04.45206Z","shell.execute_reply.started":"2023-05-14T13:56:56.196089Z","shell.execute_reply":"2023-05-14T13:57:04.450824Z"},"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])\n# 创建一个与原始图像相同大小的全白图像\ncanvas = np.ones_like(image_array) * 255\n# 将所有掩码绘制在画布上，使用不同的颜色\nfor i, mask in enumerate(masks):\n    # 转换掩码为 RGB 格式\n    real_mask = mask['segmentation'].astype(np.uint8)\n    mask_rgb = cv2.cvtColor(real_mask, cv2.COLOR_GRAY2RGB)\n    # 使用不同的颜色\n    color = (np.random.randint(0, 256), np.random.randint(0, 256), np.random.randint(0, 256))\n    # 将掩码绘制在画布上\n    canvas = cv2.addWeighted(canvas, 1, canvas, 1, 0)\n    canvas[real_mask > 0] = color\n\n# 将画布转换为 PIL 图像，并保存\nimage_pil = Image.fromarray(cv2.cvtColor(canvas, cv2.COLOR_BGR2RGB))\nimage_pil.save(\"/kaggle/working/ces.png\")","metadata":{"execution":{"iopub.status.busy":"2024-03-27T14:34:11.056969Z","iopub.execute_input":"2024-03-27T14:34:11.057877Z","iopub.status.idle":"2024-03-27T14:34:11.149308Z","shell.execute_reply.started":"2024-03-27T14:34:11.057832Z","shell.execute_reply":"2024-03-27T14:34:11.146969Z"},"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])\n# 创建一个与原始图像相同大小的全白图像\ncanvas = np.ones_like(image_array) * 255\n# 将所有掩码绘制在画布上，使用不同的颜色\nfor i, mask in enumerate(masks):\n    # 转换掩码为 RGB 格式\n    real_mask = mask['segmentation'].astype(np.uint8)\n    mask_rgb = cv2.cvtColor(real_mask, cv2.COLOR_GRAY2RGB)\n    # 使用不同的颜色\n    color = (np.random.randint(0, 256), np.random.randint(0, 256), np.random.randint(0, 256))\n    # 将掩码绘制在画布上\n    canvas = cv2.addWeighted(canvas, 1, canvas, 1, 0)\n    canvas[real_mask > 0] = color\n\n# 将画布转换为 PIL 图像，并保存\nimage_pil = Image.fromarray(cv2.cvtColor(canvas, cv2.COLOR_BGR2RGB))\nimage_pil.save(r\"D:\\A_WD14\\wd14-tagger-standalone\\ImageAssistant\\ces.png\")","metadata":{"execution":{"iopub.status.busy":"2023-05-14T13:57:14.867617Z","iopub.execute_input":"2023-05-14T13:57:14.868952Z","iopub.status.idle":"2023-05-14T13:57:22.715798Z","shell.execute_reply.started":"2023-05-14T13:57:14.868909Z","shell.execute_reply":"2023-05-14T13:57:22.714732Z"},"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-05-14T13:57:22.71748Z","iopub.execute_input":"2023-05-14T13:57:22.718592Z","iopub.status.idle":"2023-05-14T13:58:19.400307Z","shell.execute_reply.started":"2023-05-14T13:57:22.718552Z","shell.execute_reply":"2023-05-14T13:58:19.399153Z"},"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-05-14T13:58:19.402102Z","iopub.execute_input":"2023-05-14T13:58:19.402514Z","iopub.status.idle":"2023-05-14T13:58:58.136222Z","shell.execute_reply.started":"2023-05-14T13:58:19.402475Z","shell.execute_reply":"2023-05-14T13:58:58.135177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del(mask_generator)","metadata":{"execution":{"iopub.status.busy":"2023-05-14T13:58:58.137803Z","iopub.execute_input":"2023-05-14T13:58:58.138428Z","iopub.status.idle":"2023-05-14T13:58:58.143567Z","shell.execute_reply.started":"2023-05-14T13:58:58.138386Z","shell.execute_reply":"2023-05-14T13:58:58.142331Z"},"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-05-14T13:58:58.147884Z","iopub.execute_input":"2023-05-14T13:58:58.148355Z","iopub.status.idle":"2023-05-14T13:58:58.15423Z","shell.execute_reply.started":"2023-05-14T13:58:58.148318Z","shell.execute_reply":"2023-05-14T13:58:58.153098Z"},"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-05-14T13:58:58.155874Z","iopub.execute_input":"2023-05-14T13:58:58.157273Z","iopub.status.idle":"2023-05-14T13:58:59.635284Z","shell.execute_reply.started":"2023-05-14T13:58:58.15722Z","shell.execute_reply":"2023-05-14T13:58:59.634284Z"},"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-05-14T13:58:59.636961Z","iopub.execute_input":"2023-05-14T13:58:59.637635Z","iopub.status.idle":"2023-05-14T13:59:01.342742Z","shell.execute_reply.started":"2023-05-14T13:58:59.637593Z","shell.execute_reply":"2023-05-14T13:59:01.341745Z"},"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-05-14T13:59:01.344241Z","iopub.execute_input":"2023-05-14T13:59:01.3455Z","iopub.status.idle":"2023-05-14T13:59:04.187225Z","shell.execute_reply.started":"2023-05-14T13:59:01.345456Z","shell.execute_reply":"2023-05-14T13:59:04.186108Z"},"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-05-14T13:59:04.188899Z","iopub.execute_input":"2023-05-14T13:59:04.189392Z","iopub.status.idle":"2023-05-14T13:59:06.686708Z","shell.execute_reply.started":"2023-05-14T13:59:04.189351Z","shell.execute_reply":"2023-05-14T13:59:06.685756Z"},"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":{}}]}