{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n        break\n    break\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-17T20:40:25.447986Z","iopub.execute_input":"2023-06-17T20:40:25.448478Z","iopub.status.idle":"2023-06-17T20:40:25.457830Z","shell.execute_reply.started":"2023-06-17T20:40:25.448443Z","shell.execute_reply":"2023-06-17T20:40:25.456851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Install dependencies","metadata":{}},{"cell_type":"code","source":"pip install pycocotools","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:40:26.748477Z","iopub.execute_input":"2023-06-17T20:40:26.748828Z","iopub.status.idle":"2023-06-17T20:41:01.287470Z","shell.execute_reply.started":"2023-06-17T20:40:26.748797Z","shell.execute_reply":"2023-06-17T20:41:01.286311Z"},"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-06-17T20:41:01.289976Z","iopub.execute_input":"2023-06-17T20:41:01.290358Z","iopub.status.idle":"2023-06-17T20:41:16.316165Z","shell.execute_reply.started":"2023-06-17T20:41:01.290321Z","shell.execute_reply":"2023-06-17T20:41:16.314953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport cv2\nimport sys\nsys.path.append(\"..\")","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:41:16.317697Z","iopub.execute_input":"2023-06-17T20:41:16.317999Z","iopub.status.idle":"2023-06-17T20:41:19.828458Z","shell.execute_reply.started":"2023-06-17T20:41:16.317968Z","shell.execute_reply":"2023-06-17T20:41:19.827351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a = torch.tensor(0)\na.device.type == 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:41:19.830938Z","iopub.execute_input":"2023-06-17T20:41:19.831671Z","iopub.status.idle":"2023-06-17T20:41:19.851903Z","shell.execute_reply.started":"2023-06-17T20:41:19.831632Z","shell.execute_reply":"2023-06-17T20:41:19.851021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from segment_anything import sam_model_registry, SamAutomaticMaskGenerator, SamPredictor\nfrom collections import defaultdict\nfrom segment_anything.utils.transforms import ResizeLongestSide\nsam_checkpoint = \"/kaggle/input/segment-anything-models/sam_vit_h_4b8939.pth\"\nmodel_type = \"vit_h\"\n# \n# vit_b - basically does NOT work, terrible pretrained results\n# sam_checkpoint = \"/kaggle/input/segment-anything-models/sam_vit_b_01ec64.pth\"\n# model_type = \"vit_b\"\ndevice = \"cuda\"\nsam = sam_model_registry[model_type](checkpoint=sam_checkpoint)\nsam.to(device=device)\nmask_generator = SamAutomaticMaskGenerator(sam)\nmask_generator_2 = SamAutomaticMaskGenerator(\n    model=sam,\n    points_per_side=32,\n    pred_iou_thresh=0.9,\n    stability_score_thresh=0.9,\n    crop_n_layers=1,\n    crop_n_points_downscale_factor=2,\n    min_mask_region_area=120,  # Requires open-cv to run post-processing\n)\n\npredictor = SamPredictor(sam)","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:41:19.853414Z","iopub.execute_input":"2023-06-17T20:41:19.853967Z","iopub.status.idle":"2023-06-17T20:42:00.460718Z","shell.execute_reply.started":"2023-06-17T20:41:19.853935Z","shell.execute_reply":"2023-06-17T20:42:00.459710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import gc\n# torch.cuda.empty_cache()\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:42:00.462333Z","iopub.execute_input":"2023-06-17T20:42:00.462753Z","iopub.status.idle":"2023-06-17T20:42:00.468489Z","shell.execute_reply.started":"2023-06-17T20:42:00.462716Z","shell.execute_reply":"2023-06-17T20:42:00.466179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport json\nimport cv2\nimport matplotlib.pyplot as plt\nimport ipywidgets as widgets\nimport numpy as np\nimport IPython.display as ipd\n\n\ntrain = glob.glob('/kaggle/input/hubmap-hacking-the-human-vasculature/train/*')\ntest = glob.glob('/kaggle/input/hubmap-hacking-the-human-vasculature/test/*')\n\nwith open('/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl', 'r') as f:\n    polygons = [json.loads(p) for p in list(f)]\n\nimg_map = {impath.split('/')[-1].split('.')[0]: impath for impath in train}\nimg_map.update({impath.split('/')[-1].split('.')[0]: impath for impath in test})\n\npolygon_map = {polygon['id']: polygon for polygon in polygons}\n\nprint(f'total images: {len(img_map)}')\nprint(f'annotated images: {len(polygon_map)}')","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:42:00.469937Z","iopub.execute_input":"2023-06-17T20:42:00.470911Z","iopub.status.idle":"2023-06-17T20:42:05.399773Z","shell.execute_reply.started":"2023-06-17T20:42:00.470878Z","shell.execute_reply":"2023-06-17T20:42:05.398718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Comparison of SAM with point detection \nPaint the vessels and glomerulus on an image.\nWith SAM, find all masks in the image \nAlso, get all coordinates of vessels in the image and make it produce masks for those points. ","metadata":{}},{"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_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='red', marker='o', s=80, edgecolor='white', linewidth=1.25)\n    ax.scatter(neg_points[:, 0], neg_points[:, 1], color='red', marker='o', s=80, edgecolor='white', linewidth=1.25)  \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=1)\n    else:\n        color = np.array([200/255, 0/255, 0/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#---------------------------------------------------   \ndef custom_plot(title,image,specific_point):\n    input_label = np.array([1])\n    #--------------\n    masks = mask_generator_2.generate(image)\n    #--------------\n    predictor.set_image(image)\n    masks_p, scores, logits = predictor.predict(\n        point_coords=specific_point,\n        point_labels=input_label,\n        multimask_output=True,\n    )\n    #--------------\n    fig, ((ax1, ax2), (ax3, ax4)) = plt.subplots(2, 2, figsize=(15, 15))\n    #fig.suptitle(f'{title} tumor')\n    plt.axis('off')\n    ax1.imshow(image)\n    ax1.title.set_text(\"Image\")\n    ax2.imshow(image)\n    ax2.title.set_text(\"Image+Masks\")\n    show_anns(masks, ax2)\n    for i, (mask, score) in enumerate(zip(masks_p, scores)):\n        if i==0:\n            ax3.imshow(image)\n            ax3.title.set_text(\"a specific object\")\n            show_points(specific_point, input_label, ax3)\n            show_mask(mask, ax4)\n            ax4.title.set_text(f\"Mask - Score: {score:.3f}\")\n    for ax in fig.get_axes():\n        ax.label_outer()\n        ax.axis('off')\n#-----------------\ndef draw(img_id):    \n    polygon = polygon_map[img_id]\n    img = cv2.imread(img_map[img_id])\n\n    blood_vessel = 0\n    glomerulus = 0\n    unsure = 0\n    annotations = []\n    for anno in polygon['annotations']:\n        if anno['type'] == 'blood_vessel':\n            color = (0,255,0)\n            blood_vessel += 1\n            \n        elif anno['type'] == 'glomerulus':\n            color = (0,0,0)\n            glomerulus += 1\n        else:\n            color = (255,0,0)\n            unsure += 1\n\n        pts = anno['coordinates']\n        pts = np.array(pts)\n        pts = pts.reshape(-1, 1, 2)\n        annotations.append(pts)\n        cv2.polylines(img, pts, True, color, 3)\n    \n    print(f'{blood_vessel = }')\n    print(f'{glomerulus = }')\n    print(f'{unsure = }')\n\n    plt.imshow(img)\n    return annotations\n\n#----------\ndef draw_sam(img_id, annotations):\n    image = cv2.cvtColor(\n        cv2.imread('/kaggle/input/hubmap-hacking-the-human-vasculature/train/' + img_id + '.tif'), \n        cv2.COLOR_BGR2RGB\n    )\n    specific_point = np.array([[50, 100]])\n    custom_plot(img_id, image, specific_point)\n    \noutput = widgets.Output()\n\n@output.capture()\ndef ipydisplay(change):\n    img_id = change['new']\n    ipd.clear_output()\n    print(\"Drawing....\")\n    annotations = draw(img_id)\n    print(\"Finished Drawing. SAM starting...\")\n#     print(annotations)\n    draw_sam(img_id, annotations)\n    plt.axis('off')\n    plt.show()\n    \n\n# you can only use this widget when actually running the notebook\nselect = widgets.Dropdown(options=list(polygon_map.keys()))\nselect.observe(ipydisplay, 'value')\nwidgets.VBox([select, output])","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:43:45.490653Z","iopub.execute_input":"2023-06-17T20:43:45.491074Z","iopub.status.idle":"2023-06-17T20:43:45.547404Z","shell.execute_reply.started":"2023-06-17T20:43:45.491035Z","shell.execute_reply":"2023-06-17T20:43:45.546507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# np.mean(polygons[0], axis=0)","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:43:45.686771Z","iopub.execute_input":"2023-06-17T20:43:45.687217Z","iopub.status.idle":"2023-06-17T20:43:45.691644Z","shell.execute_reply.started":"2023-06-17T20:43:45.687189Z","shell.execute_reply":"2023-06-17T20:43:45.690501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_id = list(polygon_map.keys())[0]\npolygons = draw(img_id)\ncentroids = []\nfor list_of_coord in polygons: \n    centroids.append(np.mean(list_of_coord, axis=0))\n#     centroids.append(list_of_coord)","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:43:45.902874Z","iopub.execute_input":"2023-06-17T20:43:45.903761Z","iopub.status.idle":"2023-06-17T20:43:46.346486Z","shell.execute_reply.started":"2023-06-17T20:43:45.903716Z","shell.execute_reply":"2023-06-17T20:43:46.345342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def draw_original(img_id):    \n    polygon = polygon_map[img_id]\n    img = cv2.imread(img_map[img_id])\n\n    blood_vessel = 0\n    glomerulus = 0\n    unsure = 0\n    for anno in polygon['annotations']:\n        if anno['type'] == 'blood_vessel':\n            color = (0,255,0)\n            blood_vessel += 1\n        elif anno['type'] == 'glomerulus':\n            color = (0,0,0)\n            glomerulus += 1\n        else:\n            color = (255,0,0)\n            unsure += 1\n\n        pts = anno['coordinates']\n        pts = np.array(pts)\n        pts = pts.reshape(-1, 1, 2)\n        cv2.polylines(img, pts, True, color, 3)\n    \n    print(f'{blood_vessel = }')\n    print(f'{glomerulus = }')\n    print(f'{unsure = }')\n\n    return img\n\ndef custom_plot_whole_image(title, image, points_labels):\n    input_label = np.array([1])\n\n    # Generate masks for the entire image\n    masks = mask_generator_2.generate(image)\n    predictor.set_image(image)\n    \n    # Create a placeholder for the combined masks\n    combined_masks = np.zeros_like(image)\n\n    fig, axs = plt.subplots(1, 4, figsize=(15, 15))\n\n    axs[0].imshow(image)\n    axs[0].title.set_text(\"Image\")\n    centroid_masks = []\n    \n    for i, coord in enumerate(points_labels):\n        print(\"Computing mask for point \", i+1, \"/\", len(points_labels))\n        # For each point, predict and get the mask\n        masks_p, scores, logits = predictor.predict(\n            point_coords=coord,\n            point_labels=input_label,\n            multimask_output=True,\n        )\n        # Binarize and convert to uint8\n        binary_mask = (masks_p[0] > 0).astype(np.uint8)\n        centroid_masks.append(masks_p[0])\n        # Extract the border of the mask\n        dilated_mask = cv2.dilate(binary_mask, np.ones((3,3), np.uint8), iterations=1)\n        border_mask = dilated_mask - binary_mask\n        border_mask = np.repeat(border_mask[:, :, np.newaxis], 3, axis=2)\n        # Add the border to the combined_masks array\n        combined_masks += border_mask\n\n    # Normalize the combined_masks array to [0, 1]\n    combined_masks = combined_masks / combined_masks.max()\n\n    # Plotting\n    axs[1].imshow(image)\n    axs[1].title.set_text(\"Image+Masks\")\n    show_anns(masks, axs[1])\n    axs[2].imshow(combined_masks, cmap='jet', alpha=0.5)\n    axs[2].title.set_text(\"All Masks Combined\")\n    axs[3].imshow(draw_original(title))\n    axs[3].title.set_text(\"Labeled\")\n\n    for ax in fig.get_axes(): \n        ax.label_outer()\n        ax.axis('off')\n        \n    return centroid_masks\n\nimage = cv2.cvtColor(\n    cv2.imread('/kaggle/input/hubmap-hacking-the-human-vasculature/train/' + img_id + '.tif'), \n    cv2.COLOR_BGR2RGB\n)\ncentroid_masks = custom_plot_whole_image(img_id, image, centroids)\n# custom_plot_whole_image(imgid, image, annotations[0])","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:43:46.347923Z","iopub.execute_input":"2023-06-17T20:43:46.348251Z","iopub.status.idle":"2023-06-17T20:44:16.244587Z","shell.execute_reply.started":"2023-06-17T20:43:46.348221Z","shell.execute_reply":"2023-06-17T20:44:16.243689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pycocotools import mask as mask_utils\n    \ndef convert_polygon_to_binary(polygon):\n    polygon = np.array(polygon)\n    polygon = polygon.reshape((-1, 1, 2))\n    # Create an empty mask\n    mask = np.zeros((512, 512), dtype=np.uint8)\n    # Fill the polygon area in the mask\n    cv2.fillPoly(mask, [polygon], color=1)\n    return mask\n\ndef rle_to_polygon(rle):\n    # Assume rle is your RLE-encoded mask\n    binary_mask = mask_utils.decode(rle)\n    # Find contours\n    contours, _ = cv2.findContours(binary_mask, cv2.RETR_TREE, cv2.CHAIN_APPROX_SIMPLE)\n    return contours\n\ndef iou_mask_and_polygon(mask, polygon, original_img=False):\n    binary_mask = (mask > 0).astype(np.uint8)\n    polygon_in_rle = convert_polygon_to_binary(polygon)\n    rle1 = mask_utils.encode(np.asfortranarray(binary_mask))\n    rle2 = mask_utils.encode(np.asfortranarray(polygon_in_rle))\n    iou = mask_utils.iou([rle1], [rle2], [False])\n    \n    if original_img:\n        original_img = cv2.imread(img_map[img_id])\n\n        # Decoding RLE to binary mask for display\n        binary_mask_polygon = mask_utils.decode(rle2)\n\n        # Convert binary masks to RGB for visualization\n        mask_rgb = cv2.cvtColor(binary_mask, cv2.COLOR_GRAY2RGB)\n        mask_rgb_polygon = cv2.cvtColor(binary_mask_polygon, cv2.COLOR_GRAY2RGB)\n\n        # Colorize the masks\n        mask_rgb[..., 1] = 0  # zero out the green channel\n        mask_rgb[..., 2] = 0  # zero out the blue channel\n        mask_rgb = mask_rgb * 255  # scale mask to [0, 255]\n\n        mask_rgb_polygon[..., 0] = 0  # zero out the red channel\n        mask_rgb_polygon[..., 1] = 0  # zero out the green channel\n        mask_rgb_polygon = mask_rgb_polygon * 255  # scale mask to [0, 255]\n\n        # Overlay the masks on the original image\n        overlayed_img = cv2.addWeighted(original_img, 0.5, mask_rgb, 0.5, 0)\n        overlayed_img_polygon = cv2.addWeighted(original_img, 0.5, mask_rgb_polygon, 0.5, 0)\n\n        # Plotting\n        fig, axs = plt.subplots(1, 3, figsize=(15, 5))\n\n        axs[0].imshow(original_img)\n        axs[0].set_title(\"Original Image\")\n        axs[0].axis('off')\n\n        axs[1].imshow(overlayed_img)\n        axs[1].set_title(\"SAM Pretrained, IoU:\")\n        axs[1].axis('off')\n\n        axs[2].imshow(overlayed_img_polygon)\n        axs[2].set_title(\"Ground Truth\")\n        axs[2].axis('off')\n\n        plt.show()\n    return iou","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:44:16.246516Z","iopub.execute_input":"2023-06-17T20:44:16.246856Z","iopub.status.idle":"2023-06-17T20:44:16.264424Z","shell.execute_reply.started":"2023-06-17T20:44:16.246825Z","shell.execute_reply":"2023-06-17T20:44:16.263421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"polygon = polygon_map[img_id]\n\nfor i, centroid_mask in enumerate(centroid_masks):\n    iou = iou_mask_and_polygon(\n        centroid_mask, \n        polygon['annotations'][i]['coordinates'],\n        original_img=img_map[img_id]\n    )\n    print(iou)","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:44:16.265718Z","iopub.execute_input":"2023-06-17T20:44:16.266199Z","iopub.status.idle":"2023-06-17T20:44:21.334121Z","shell.execute_reply.started":"2023-06-17T20:44:16.266164Z","shell.execute_reply":"2023-06-17T20:44:21.333275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Components of SAM\nTask: Promptable Segmentation\n  - Given a: \n      - Segmentation prompt, such as a bounding box, a point, or another mask\n      - Image\n  - Return a mask\n\nModel: Segment Anything Model\n- prompt -> prompt Encoder -----------\n                                   |\n                                   v\n- image  -> image Encoder -> lightweightMaskDecoder -> Mask\n\nWhat do we need to finetune if we want to detect capillaries?\nOne possible way is to consider training samples which are composed of:\n  - <image, annotation type, mask>\n\nWe can train on these pairs, and minimize the loss of the mask produced given an annotation type, and a mask.\nThere's a bit of a problem though - how to handle multiple annotation types...\n\n### Inference \n\nDuring inference, you are given:\n  - <image, annotation type='capillary'> -> return mask\n\nHopefully, learning during training to find glomerulus will help distinguishing with capillaries?\nOr is it better to formulate this as a binary classification problem of capillaries? \n\n### Which parts can you finetune\n\n* image encoder - assumption: the image encoder must be great for extracting features. 256x64x64 output seems to work zero-shot for any kind of mask. Probably not worth finetuning, nothing seems to indicate to me that capillaries need a different kind of feature than other masks.\n* prompt encoder\n* mask decoder\n\n\n## Fine-tuning\n\n\n* Each example consists of:\n  - Pixel values (which is the image)\n  - A prompt, which is: a bounding box of the whole image? a text?\n  - A ground truth segmentation mask\n  ","metadata":{}},{"cell_type":"markdown","source":"## Dataset\n\nThe dataset class takes a list of images, and a list of masks. \nThe masks ","metadata":{}},{"cell_type":"code","source":"sam_checkpoint","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:44:21.336712Z","iopub.execute_input":"2023-06-17T20:44:21.337811Z","iopub.status.idle":"2023-06-17T20:44:21.344126Z","shell.execute_reply.started":"2023-06-17T20:44:21.337773Z","shell.execute_reply":"2023-06-17T20:44:21.343293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data preprocessing \n\nNow it's time to convert the input images into a format SAM's image encoder expects. We can use the internal tools `ResizeLongestSide, sam.preprocess` etc.. for this.\n\nProblems now: \n - Need to include the prompts (type of prediction) and ground truth (masks) into this solution.\n - How to use CLIP embeddings as prompts? https://github.com/suvansh/say-anything-backend ","metadata":{}},{"cell_type":"code","source":"def get_bounding_box(ground_truth_map):\n    # get bounding box from mask\n    y_indices, x_indices = np.where(ground_truth_map > 0)\n    x_min, x_max = np.min(x_indices), np.max(x_indices)\n    y_min, y_max = np.min(y_indices), np.max(y_indices)\n    # add perturbation to bounding box coordinates\n    H, W = ground_truth_map.shape\n    x_min = max(0, x_min - np.random.randint(0, 1))\n    x_max = min(W, x_max + np.random.randint(0, 1))\n    y_min = max(0, y_min - np.random.randint(0, 1))\n    y_max = min(H, y_max + np.random.randint(0, 1))\n    bbox = np.array([x_min, y_min, x_max, y_max])\n\n    return bbox\n\ndef get_image_mask_and_bbox(img_ids):\n    # For every bounding box(prompt), find the original image and ground truth mask\n    bbox_coords = defaultdict(list)\n    ground_truth_masks = defaultdict(list)\n    transformed_data = defaultdict(dict)\n    \n    for img_id in img_ids: \n        image = cv2.imread(img_map[img_id])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)  \n        transform = ResizeLongestSide(sam.image_encoder.img_size)\n        input_image = transform.apply_image(image)\n        input_image_torch = torch.as_tensor(input_image, device=device)\n        transformed_image = input_image_torch.permute(2, 0, 1).contiguous()[None, :, :, :]  \n\n        input_image = sam.preprocess(transformed_image)\n        original_image_size = image.shape[:2]\n        input_size = tuple(transformed_image.shape[-2:])\n\n        transformed_data[img_id]['image'] = input_image[0]\n        transformed_data[img_id]['input_size'] = input_size\n        transformed_data[img_id]['original_image_size'] = original_image_size\n\n        for annotation in polygon_map[img_id]['annotations']:\n            if annotation['type'] == 'blood_vessel':\n                ground_truth_mask = convert_polygon_to_binary(annotation['coordinates'])\n                ground_truth_masks[img_id].append(ground_truth_mask)\n                bbox_coords[img_id].append(get_bounding_box(ground_truth_mask))\n                # Turn these coordinates into a ground truth mask, and a bounding box prompt\n                # plt.imshow(ground_truth_mask)\n    return transformed_data, ground_truth_masks, bbox_coords","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:44:21.345801Z","iopub.execute_input":"2023-06-17T20:44:21.346453Z","iopub.status.idle":"2023-06-17T20:44:21.364049Z","shell.execute_reply.started":"2023-06-17T20:44:21.346421Z","shell.execute_reply":"2023-06-17T20:44:21.362959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BloodVesselDataset(torch.utils.data.Dataset):\n    def __init__(self, img_ids, device):\n        self.img_ids = img_ids\n        self.device = device\n        results = get_image_mask_and_bbox(img_ids)\n        self.transformed_data, self.ground_truth_masks, self.bbox_coords = results \n        \n    def __len__(self):\n        return len(self.img_ids)\n\n    def __getitem__(self, idx):\n        img_id = self.img_ids[idx]\n        return self.transformed_data[img_id], self.ground_truth_masks[img_id], self.bbox_coords[img_id]\n\ntrain_len, test_len = 200, 20\ntrain_dataset = BloodVesselDataset(list(polygon_map.keys())[:train_len], device)\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=1, shuffle=True) \ntest_dataset = BloodVesselDataset(list(polygon_map.keys())[:test_len], device)\nval_loader = torch.utils.data.DataLoader(test_dataset, batch_size=1, shuffle=True) ","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:44:21.365591Z","iopub.execute_input":"2023-06-17T20:44:21.366323Z","iopub.status.idle":"2023-06-17T20:44:31.987630Z","shell.execute_reply.started":"2023-06-17T20:44:21.366288Z","shell.execute_reply":"2023-06-17T20:44:31.986641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from segment_anything import sam_model_registry, SamAutomaticMaskGenerator, SamPredictor\nfrom collections import defaultdict\nfrom segment_anything.utils.transforms import ResizeLongestSide\nsam_checkpoint = \"/kaggle/input/segment-anything-models/sam_vit_h_4b8939.pth\"\nmodel_type = \"vit_h\"\ndevice = \"cuda\"\nsam = sam_model_registry[model_type](checkpoint=sam_checkpoint)\nsam.to(device=device)\n\n# Set up the optimizer, hyperparameter tuning will improve performance here\n# lr = 1e-3\n# wd = 0\n# optimizer = torch.optim.Adam(sam.mask_decoder.parameters(), lr=lr, weight_decay=wd)\n# scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10, eta_min=0)\n# loss_fn = torch.nn.MSELoss()\n\n# Consider changing the loss function to Dice Loss or BCE Loss\n# class DiceLoss(torch.nn.Module):\n#     def __init__(self, eps=1e-5):\n#         super().__init__()\n#         self.eps = eps\n        \n#     def forward(self, output, target):\n#         intersection = (output * target).sum()\n#         union = output.sum() + target.sum() + self.eps\n#         dice_score = 2.0 * intersection / union\n#         return -(1.0 - dice_score)\n\n# loss_fn = DiceLoss()\nloss_fn = torch.nn.BCELoss()","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:44:31.989148Z","iopub.execute_input":"2023-06-17T20:44:31.989524Z","iopub.status.idle":"2023-06-17T20:44:59.951588Z","shell.execute_reply.started":"2023-06-17T20:44:31.989490Z","shell.execute_reply":"2023-06-17T20:44:59.950548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset validation check","metadata":{}},{"cell_type":"code","source":"# from https://www.kaggle.com/code/paulorzp/rle-functions-run-lenght-encode-decode\ndef rle2mask(mask_rle, shape=(1600,256)):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (width,height) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T\n\ndef mask2rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels= img.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\n# |mask2rle(ground_truth_masks[0][0])","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:47:53.685270Z","iopub.execute_input":"2023-06-17T20:47:53.685669Z","iopub.status.idle":"2023-06-17T20:47:53.695570Z","shell.execute_reply.started":"2023-06-17T20:47:53.685639Z","shell.execute_reply":"2023-06-17T20:47:53.694391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.patches as patches\n\n# Assuming that ground_truth_masks and bbox_coords are dicts indexed by image id\ni = 0\nfor transformed_data, ground_truth_masks, bbox_coords in train_loader:\n    i += 1\n    if i > 1:\n        break\n        \n    print(f\"Processing image {i}/{len(train_loader)} with {len(ground_truth_masks)} blood vessels\") \n\n    # Get all masks for capillaries in a single img\n    input_image = transformed_data['image'].to(device)\n\n    fig, ax = plt.subplots(1, 1, figsize=(10, 10))\n    input_image_np = input_image[0].cpu().numpy().transpose(1,2,0)\n    input_image_np = cv2.resize(input_image_np, (512, 512))\n    masks_canvas = np.zeros(input_image_np.shape)\n\n    for i in range(len(ground_truth_masks)):\n        mask = ground_truth_masks[i][0]\n        # Create a colored mask\n        colored_mask = np.zeros((*mask.shape, 3))  # create an array of zeros with the same height and width as the mask, but with an extra dimension for color channels\n        colored_mask[mask > 0] = [0, 1, 0]  # make the mask green\n\n        # Overlay the mask onto the image\n        masks_canvas = np.maximum(masks_canvas, colored_mask)  # keep the maximum value at each pixel (i.e., if a pixel is covered by multiple masks, it keeps the\n        # Draw the bounding box\n        bbox = bbox_coords[i][0]  # assuming bbox is a list/tuple in the format (xmin, ymin, xmax, ymax)\n        rect = patches.Rectangle((bbox[0], bbox[1]), bbox[2] - bbox[0], bbox[3] - bbox[1], linewidth=1, edgecolor='r', facecolor='none')\n        ax.add_patch(rect)\n        \n    overlay = (input_image_np * 0.25) + (masks_canvas * 0.75)  # adjust weights as needed\n    ax.imshow(overlay)","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:47:54.468513Z","iopub.execute_input":"2023-06-17T20:47:54.468889Z","iopub.status.idle":"2023-06-17T20:47:55.499344Z","shell.execute_reply.started":"2023-06-17T20:47:54.468857Z","shell.execute_reply":"2023-06-17T20:47:55.493565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\nfrom pytorch_lightning.loggers import WandbLogger\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"api_wandb\")\n\nepochs = 30\nwandb.login(key=wandb_api)\nwandb.init(\n    # set the wandb project where this run will be logged\n    project=\"Capillaries\",\n    # track hyperparameters and run metadata\n    config={\n    \"learning_rate\": 1e-3,\n    \"architecture\": \"SAM-finetune mask decoder\",\n    \"dataset\": \"Only 20 images\",\n    \"epochs\": epochs,\n    }\n)\nwandb_logger = WandbLogger()","metadata":{"execution":{"iopub.status.busy":"2023-06-17T20:45:00.411977Z","iopub.status.idle":"2023-06-17T20:45:00.412809Z","shell.execute_reply.started":"2023-06-17T20:45:00.412571Z","shell.execute_reply":"2023-06-17T20:45:00.412594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom torch.nn.functional import threshold, normalize\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\nfrom torchvision.utils import make_grid\n\nclass MyModel(pl.LightningModule):\n    def __init__(self, sam, loss_fn, plot_every_x_steps):\n        super(MyModel, self).__init__()\n        self.sam = sam\n        self.loss_fn = loss_fn\n        self.plot_every_x_steps = plot_every_x_steps\n        self.transform = ResizeLongestSide(sam.image_encoder.img_size)\n    \n    def forward_pass(self, transformed_data, ground_truth_masks, bbox_coords):\n        input_image = transformed_data['image'].to(self.device)\n        input_size = transformed_data[\"input_size\"]\n        original_image_size = transformed_data[\"original_image_size\"]\n        input_size = (input_size[0].item(), input_size[1].item())\n        original_image_size = (original_image_size[0].item(), original_image_size[1].item())\n        losses = []\n\n        for i in range(len(ground_truth_masks)):\n            ground_truth_masks[i] = ground_truth_masks[i][0].to(self.device)\n            prompt_box = bbox_coords[i][0].cpu().numpy()\n\n            with torch.no_grad():\n                image_embedding = self.sam.image_encoder(input_image)\n                box = self.transform.apply_boxes(prompt_box, original_image_size)\n                box_torch = torch.as_tensor(box, dtype=torch.float, device=self.device)\n                box_torch = box_torch[None, :]\n\n                sparse_embeddings, dense_embeddings = self.sam.prompt_encoder(\n                    points=None,\n                    boxes=box_torch,\n                    masks=None,\n                )\n\n            low_res_masks, _ = self.sam.mask_decoder(\n                image_embeddings=image_embedding,\n                image_pe=self.sam.prompt_encoder.get_dense_pe(),\n                sparse_prompt_embeddings=sparse_embeddings,\n                dense_prompt_embeddings=dense_embeddings,\n                multimask_output=False,\n            )\n\n            upscaled_masks = self.sam.postprocess_masks(low_res_masks, input_size, original_image_size).to(self.device)\n            probability_mask = torch.sigmoid(upscaled_masks[0,0])\n            loss = self.loss_fn(probability_mask, ground_truth_masks[i].float())\n            losses.append(loss)\n            last_upscaled_mask = upscaled_masks\n\n        mean_loss = torch.stack(losses).mean()\n\n        return mean_loss, last_upscaled_mask\n\n    def training_step(self, batch, batch_idx):\n        transformed_data, ground_truth_masks, bbox_coords = batch\n        if len(ground_truth_masks) == 0:\n            return None\n        self.last_imgs = transformed_data['image'] \n        self.last_masks = ground_truth_masks[-1]\n        mean_loss, last_upscaled_mask = self.forward_pass(transformed_data, ground_truth_masks, bbox_coords)\n        self.last_upscaled_masks = last_upscaled_mask\n        self.log('train_loss', mean_loss)\n        self.log_images()\n        return mean_loss\n\n    @torch.no_grad()\n    def validation_step(self, batch, batch_idx):\n        transformed_data, ground_truth_masks, bbox_coords = batch\n        if len(ground_truth_masks) == 0:\n            return None\n        self.last_imgs = transformed_data['image'] \n        self.last_masks = ground_truth_masks[-1]\n        mean_loss, last_upscaled_mask = self.forward_pass(transformed_data, ground_truth_masks, bbox_coords)\n        self.last_upscaled_masks = last_upscaled_mask\n        self.log('val_loss', mean_loss)\n        self.log_images()\n        return mean_loss   \n    \n    def configure_optimizers(self):\n        lr = 1e-3\n        wd = 0\n        optimizer = torch.optim.Adam(self.sam.mask_decoder.parameters(), lr=lr, weight_decay=wd)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10, eta_min=0)\n\n        return {\"optimizer\": optimizer, \"lr_scheduler\": scheduler, \"monitor\": \"train_loss\"}\n     \n    def log_images(self):\n        if True or self.global_step % self.plot_every_x_steps == 0:  # plot every x epochs\n            last_img = self.last_imgs[0].cpu().numpy().transpose(1,2,0)\n            last_mask = self.last_masks[0].detach().cpu().numpy()\n            last_upscaled_mask = self.last_upscaled_masks[0,0].detach().cpu().numpy()\n#             fig, ax = plt.subplots(3, 1, figsize=(15, 15))\n\n#             ax[0].imshow(last_img)\n#             ax[0].title.set_text('Input Images')\n\n#             ax[1].imshow(last_mask, cmap='gray')\n#             ax[1].title.set_text('Ground Truth Masks')\n\n#             ax[2].imshow(last_upscaled_mask, cmap='gray')\n#             ax[2].title.set_text('Upscaled Masks')\n\n#             plt.show()\n            self.logger.experiment.log({\n                \"Input Images\": wandb.Image(last_img),\n                \"Ground Truth Masks\": wandb.Image(last_mask, mode='L'),\n                \"Predicted (upscaled) Masks\": wandb.Image(last_upscaled_mask, mode='L'),\n            })\n        \n\nmodel = MyModel(sam, loss_fn, plot_every_x_steps=10)\ncheckpoint_callback = ModelCheckpoint(\n    dirpath=\"checkpoints\",\n    filename=\"best-checkpoint\",\n    save_top_k=1,\n    verbose=True,\n    monitor=\"val_loss\",  # assumes you have a validation step where you log a \"val_loss\"\n    mode=\"min\",\n    every_n_epochs=5,  # change this to save every X epochs\n)\nearly_stopping_callback = EarlyStopping(\n    monitor=\"val_loss\",  # assumes you have a validation step where you log a \"val_loss\"\n    patience=3,\n    mode='min'\n)\ntrainer = pl.Trainer(\n    logger=wandb_logger,\n    max_epochs=epochs,\n    callbacks=[checkpoint_callback, early_stopping_callback],\n)\ntrainer.fit(model, train_loader, val_loader)\nwandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-06-06T12:41:20.166434Z","iopub.execute_input":"2023-06-06T12:41:20.169082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from IPython.display import clear_output\n# from statistics import mean\n# from tqdm import tqdm\n# from torch.nn.functional import threshold, normalize\n\n# num_epochs = 100\n# losses = []\n\n# for epoch in range(num_epochs):\n#     print(f\"Starting epoch: {epoch}\")\n#     epoch_losses = []\n#     j = 0\n#     for transformed_data, ground_truth_masks, bbox_coords in loader:\n#         j += 1\n#         print(f\"Processing image {j}/{len(loader)} with {len(ground_truth_masks)} blood vessels\")\n#         # Get all masks for capillaries in a single img\n#         input_image = transformed_data['image'].to(device)        \n#         # Convert tensors back to Python tuples\n#         input_size = transformed_data[\"input_size\"]\n#         original_image_size = transformed_data[\"original_image_size\"]\n#         input_size = (input_size[0].item(), input_size[1].item())\n#         original_image_size = (original_image_size[0].item(), original_image_size[1].item())\n        \n#         # Try all masks?\n#         for i in range(len(ground_truth_masks)):\n#             mask = ground_truth_masks[i][0]\n#             ground_truth_masks[i] = ground_truth_masks[i][0].to(device)\n            \n#             # No grad here as we don't want to optimise the image/prompt encoders\n#             with torch.no_grad():\n#                 image_embedding = sam.image_encoder(input_image)\n#                 prompt_box = bbox_coords[i][0]\n#                 transform = ResizeLongestSide(sam.image_encoder.img_size)\n#                 box = transform.apply_boxes(prompt_box.numpy(), original_image_size)\n#                 box_torch = torch.as_tensor(box, dtype=torch.float, device=device)\n#                 box_torch = box_torch[None, :]\n\n#                 sparse_embeddings, dense_embeddings = sam.prompt_encoder(\n#                   points=None,\n#                   boxes=box_torch,\n#                   masks=None,\n#                 )\n                \n#             # The mask decoder is the only one we want to optimize\n#             low_res_masks, iou_predictions = sam.mask_decoder(\n#               image_embeddings=image_embedding,\n#               image_pe=sam.prompt_encoder.get_dense_pe(),\n#               sparse_prompt_embeddings=sparse_embeddings,\n#               dense_prompt_embeddings=dense_embeddings,\n#               multimask_output=False,\n#             )\n#             upscaled_masks = sam.postprocess_masks(low_res_masks, input_size, original_image_size).to(device)\n#             probability_mask = torch.sigmoid(upscaled_masks[0,0])\n#             loss = loss_fn(probability_mask, ground_truth_masks[i].float())\n#             if i == 0 and j == 1: \n#                 plt.subplot(1, 1, 1)\n#                 plt.imshow(ground_truth_masks[i].detach().cpu().numpy(), cmap='gray')\n#                 plt.title('Ground Truth Binary Mask')\n#                 plt.show()\n#                 plt.subplot(1, 1, 2)\n#                 plt.imshow(upscaled_masks[0,0].detach().cpu().numpy(), cmap='gray')\n#                 plt.title('Upscaled mask')\n#                 plt.show()\n#             optimizer.zero_grad()\n#             loss.backward()\n#             optimizer.step()\n#             epoch_losses.append(loss.item())\n            \n#     losses.append(epoch_losses)\n#     scheduler.step()\n#     print(f'Epoch: {epoch} finished')\n#     print(f'Mean loss: {mean(epoch_losses)}')\n    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# assuming upscaled_masks, low_res_masks are 4D tensors in the format [batch_size, channels, height, width]\n# upscaled_mask = upscaled_masks[0,0]\nbinary_mask = np.zeros((*upscaled_masks.shape, 1))\nbinary_mask[upscaled_masks.detach().cpu() > 0] = [1]\nbinary_mask = binary_mask.squeeze(-1)\nplt.imshow(upscaled_masks[0,0].detach().cpu(), cmap='gray')\nplt.title('Upscaled Masks')\nplt.show()\n# plt.imshow(upscaled_masks[0,0].detach().cpu(), cmap='gray')\n# plt.title('Upscaled Masks')\n# plt.show()\nplt.imshow(torch.sigmoid(upscaled_masks[0,0]).detach().cpu(), cmap='gray')\nplt.title('binary Masks')\nplt.show()\nplt.imshow(gt_mask.detach().cpu(), cmap='gray')\nplt.title('GT Masks')\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sam.mask_threshold","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}