{"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":"# HuBMAP - Hacking the Human Vasculature\nThe goal of this competition is to segment instances of microvascular structures, including capillaries, arterioles, and venules. The task is to create a model trained on 2D PAS-stained histology images from healthy human kidney tissue slides.","metadata":{"_uuid":"4c52ded9-bdde-4cf4-bc87-3fb4355585f9","_cell_guid":"090a1d77-c428-42ba-bfa0-e1fdb7aa9aa6","trusted":true}},{"cell_type":"markdown","source":"# Libraries","metadata":{"_uuid":"6f3a3397-a40f-429b-b4bf-02f53d4ae9c8","_cell_guid":"5578bbf8-2d9e-4f32-8b90-82e65b8da224","trusted":true}},{"cell_type":"code","source":"!pip install segmentation-models-pytorch","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-19T00:21:01.686656Z","iopub.execute_input":"2023-06-19T00:21:01.687185Z","iopub.status.idle":"2023-06-19T00:21:13.309133Z","shell.execute_reply.started":"2023-06-19T00:21:01.687139Z","shell.execute_reply":"2023-06-19T00:21:13.307843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install pipreqsnb","metadata":{"_uuid":"40d37f73-7f95-4d1d-bae2-c6cc830ba877","_cell_guid":"7143ad04-e376-4721-9d9d-6db920b32b37","collapsed":false,"jupyter":{"outputs_hidden":false},"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-19T00:21:13.311398Z","iopub.execute_input":"2023-06-19T00:21:13.311740Z","iopub.status.idle":"2023-06-19T00:21:24.962480Z","shell.execute_reply.started":"2023-06-19T00:21:13.311709Z","shell.execute_reply":"2023-06-19T00:21:24.961310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json # Json data handling utilities\nimport os # OS utilities\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport matplotlib.pyplot as plt # Plotting\n\n# ML framework\nimport torch \nfrom torch import nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom tqdm.auto import tqdm\n\n# Segmentation models and utilities\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.losses import DiceLoss\n\n# Read TIF files\nimport rasterio\nfrom rasterio.plot import show\n\n# Image augmentation \nimport albumentations as A ","metadata":{"_uuid":"9199904b-5fdc-4d04-a715-ecd89f59b7a7","_cell_guid":"7763839d-0c4c-4210-bc10-77a55539bece","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:24.964605Z","iopub.execute_input":"2023-06-19T00:21:24.965002Z","iopub.status.idle":"2023-06-19T00:21:24.974734Z","shell.execute_reply.started":"2023-06-19T00:21:24.964965Z","shell.execute_reply":"2023-06-19T00:21:24.972058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Constants","metadata":{}},{"cell_type":"code","source":"DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nIMG_SIZE = 512\n\nBATCH_SIZE = 8\nEPOCHS = 15\nLR = 1e-03\n\nENCODER = 'timm-efficientnet-b0'\nWEIGHTS = 'imagenet'\n\nMODEL_FILE = '/kaggle/working/best_hubmap_model.pt'","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:24.977971Z","iopub.execute_input":"2023-06-19T00:21:24.978317Z","iopub.status.idle":"2023-06-19T00:21:24.987508Z","shell.execute_reply.started":"2023-06-19T00:21:24.978286Z","shell.execute_reply":"2023-06-19T00:21:24.986562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data paths","metadata":{"_uuid":"11881b10-9652-4e6f-97e1-a0ff7b167eb6","_cell_guid":"9c87475f-3ae9-4f66-bc64-0b6a324cc891","trusted":true}},{"cell_type":"code","source":"BASE_DIR = '/kaggle/input/hubmap-hacking-the-human-vasculature/'\nTRAIN_DIR = os.path.join(BASE_DIR, 'train')\nTEST_DIR = os.path.join(BASE_DIR, 'test')\nPOLYGONS_FILE = os.path.join(BASE_DIR, 'polygons.jsonl')\nTILE_META_FILE = os.path.join(BASE_DIR, 'tile_meta.csv')\nWSI_META_FILE = os.path.join(BASE_DIR, 'wsi_meta.csv')","metadata":{"_uuid":"496d9f26-da9d-4c3d-a348-5157f37e1a49","_cell_guid":"65d57bd4-5aa0-4ca0-8c23-69b36fc457aa","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:24.989191Z","iopub.execute_input":"2023-06-19T00:21:24.989613Z","iopub.status.idle":"2023-06-19T00:21:25.004979Z","shell.execute_reply.started":"2023-06-19T00:21:24.989580Z","shell.execute_reply":"2023-06-19T00:21:25.004075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Both, ```TRAIN_DIR``` and ```TEST_DIR``` are folders containing TIFF images of the tiles. Each tile is 512x512 in size.\n\n```POLYGONS_FILE``` is a file containing polygonal segmentation masks in JSONL format, available for Dataset 1 and Dataset 2.\n\n```TILE_META_FILE``` is a file containing metadata for each image.\n\n```WSI_META_FILE``` is a file containing metadata for the Whole Slide Images the tiles were extracted from.\n\nMore detailed information is available at [competition data page](https://www.kaggle.com/competitions/hubmap-hacking-the-human-vasculature/data).","metadata":{"_uuid":"350022e1-9cf0-4662-8b0e-c013f41335ad","_cell_guid":"6bf0f8d8-fc9d-46c1-a2d5-3cf30aaa0269","trusted":true}},{"cell_type":"markdown","source":"# EDA","metadata":{"_uuid":"72566e20-f9b9-425d-afe5-a0468cc11967","_cell_guid":"10a5c4a9-2716-45b2-93ca-1db186bfe299","trusted":true}},{"cell_type":"markdown","source":"Before doing preprocess and prepare our data, I will like to explore the data we are going to work with.","metadata":{"_uuid":"527b8beb-dbe1-48df-8bc8-abdd09787131","_cell_guid":"dec4827d-16d8-4cb0-b42e-41d7dad7edcc","trusted":true}},{"cell_type":"code","source":"wsi_meta = pd.read_csv(WSI_META_FILE)\nwsi_meta","metadata":{"_uuid":"9e5ea774-5b28-4517-b661-3601375d0771","_cell_guid":"2218c12b-b3e1-4058-b23e-d46e2a5a2847","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:25.006536Z","iopub.execute_input":"2023-06-19T00:21:25.007308Z","iopub.status.idle":"2023-06-19T00:21:25.032654Z","shell.execute_reply.started":"2023-06-19T00:21:25.007245Z","shell.execute_reply":"2023-06-19T00:21:25.031710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_meta = pd.read_csv(TILE_META_FILE)\n\nprint('Number of rows')\nprint(len(tile_meta.index))\n\nprint('\\nFirst five rows')\nprint(tile_meta.head(5))","metadata":{"_uuid":"7279b7e0-0a8d-4272-859a-410ae7aa91a5","_cell_guid":"77b18fa9-6325-4235-a4ea-a8e0a62fb820","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:25.034017Z","iopub.execute_input":"2023-06-19T00:21:25.034474Z","iopub.status.idle":"2023-06-19T00:21:25.064197Z","shell.execute_reply.started":"2023-06-19T00:21:25.034438Z","shell.execute_reply":"2023-06-19T00:21:25.063089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Statistics')\ntile_meta.describe()","metadata":{"_uuid":"53853cd9-24a6-466f-9813-a43675f6e67d","_cell_guid":"f551f290-d3d7-4a38-909a-5fe1c07521ff","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:25.066050Z","iopub.execute_input":"2023-06-19T00:21:25.066451Z","iopub.status.idle":"2023-06-19T00:21:25.093616Z","shell.execute_reply.started":"2023-06-19T00:21:25.066415Z","shell.execute_reply":"2023-06-19T00:21:25.092472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(POLYGONS_FILE, 'r') as json_file:\n    annotations = [json.loads(line) for line in json_file]","metadata":{"_uuid":"1c65ee8d-848f-41e5-8974-b2dcbb042437","_cell_guid":"67f57c73-97f5-4dca-a002-9dade55af79b","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:25.095293Z","iopub.execute_input":"2023-06-19T00:21:25.095651Z","iopub.status.idle":"2023-06-19T00:21:29.229003Z","shell.execute_reply.started":"2023-06-19T00:21:25.095620Z","shell.execute_reply":"2023-06-19T00:21:29.227797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(annotations)","metadata":{"_uuid":"75d442cc-418c-4b91-9210-ee728b04b137","_cell_guid":"a484b0d6-b22c-4458-85fe-288dcfbf59ef","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:29.233713Z","iopub.execute_input":"2023-06-19T00:21:29.234095Z","iopub.status.idle":"2023-06-19T00:21:29.240178Z","shell.execute_reply.started":"2023-06-19T00:21:29.234049Z","shell.execute_reply":"2023-06-19T00:21:29.239155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As we can see, while our train directory contains 7033 images, only 1633 of them are annotated, so our training data is reduced from 7033 to 1633.","metadata":{"_uuid":"d2a43416-82ea-4304-9072-adeb70bbebfb","_cell_guid":"732adb82-68eb-4b14-929a-25c06c4dd612","trusted":true}},{"cell_type":"markdown","source":"# Custom PyTorch Dataset and DataLoader\nIn order to join the images with their corresponding annotations we are going to create a custom dataset using Dataset PyTorch class, this will also be useful to create a DataLoader.","metadata":{"_uuid":"aac50d1f-c53b-4e2c-8a82-2b8f9d298b54","_cell_guid":"c347a682-259c-417d-830d-32854a42dc39","trusted":true}},{"cell_type":"markdown","source":"## Augmentation function","metadata":{}},{"cell_type":"code","source":"def get_train_augs():\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n    ])","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:29.241683Z","iopub.execute_input":"2023-06-19T00:21:29.242307Z","iopub.status.idle":"2023-06-19T00:21:29.253451Z","shell.execute_reply.started":"2023-06-19T00:21:29.242264Z","shell.execute_reply":"2023-06-19T00:21:29.252368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubmapDataset(Dataset):\n    def __init__(self, img_dir, annotations_file, augmentations=None):\n        \"\"\"\n        Parameters\n        img_dir: Image directory (os.path)\n        annotations_file: File that stores the annotations of each image (os.path)\n        augmentations: Augmentations operations to apply to image.\n        \n        Return\n        image: Image read as PyTorch Tensor \n        mask: Mask read as PyTorch Tensor\n        \"\"\"\n        with open(annotations_file, 'r') as json_file:\n            self.annotations = [json.loads(line) for line in json_file]\n            \n        self.img_dir = img_dir\n        self.augmentations = augmentations\n\n    def __len__(self):\n        return len(self.annotations)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_dir, f\"{self.annotations[idx]['id']}.tif\")\n        \n        with rasterio.open(img_path) as image:\n            # Shape: [C, H, W]\n            image_array = image.read()                                            \n                    \n        # Initialize mask\n        mask = np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.float32)\n        \n        for annot in annotations[idx]['annotations']:\n            if annot['type'] == 'blood_vessel':\n                for cord in annot['coordinates']:\n                    # Note: Here height value is in position 1\n                    #       and width value is in position 0\n                    h = [i[1] for i in cord]\n                    w = [i[0] for i in cord]\n                \n                    mask[h, w] = 1    \n                    \n        # (H, W) -> (1, H, W)\n        mask = np.reshape(mask, [1, *mask.shape])\n        \n        if self.augmentations is not None:\n            # Note: image_array is already a numpy array.\n            data = self.augmentations(image=image_array, mask=mask)\n            image = data['image']\n            mask = data['mask']\n        \n        # Scaling\n        image = torch.from_numpy(image.copy()) / 255.0\n        mask = torch.from_numpy(mask.copy())\n\n        return image, mask","metadata":{"_uuid":"63db8d79-f07b-4066-b09e-125e5b0c503c","_cell_guid":"907b59fa-5578-4555-904c-4a5670bb5097","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:29.254898Z","iopub.execute_input":"2023-06-19T00:21:29.255391Z","iopub.status.idle":"2023-06-19T00:21:29.268178Z","shell.execute_reply.started":"2023-06-19T00:21:29.255355Z","shell.execute_reply":"2023-06-19T00:21:29.267319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's test our HubmapDataset class.","metadata":{"_uuid":"d841f3bf-1b5d-4995-9554-122c9462356a","_cell_guid":"bf732b8f-0640-48c0-8ecd-8087d682c43f","trusted":true}},{"cell_type":"code","source":"dataset = HubmapDataset(TRAIN_DIR, POLYGONS_FILE, get_train_augs())","metadata":{"_uuid":"3c2bc9fd-6554-4577-8c33-a72c029e80a8","_cell_guid":"f10a5d4a-5cfd-429d-b058-b43539217934","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:29.269600Z","iopub.execute_input":"2023-06-19T00:21:29.270039Z","iopub.status.idle":"2023-06-19T00:21:32.964003Z","shell.execute_reply.started":"2023-06-19T00:21:29.270007Z","shell.execute_reply":"2023-06-19T00:21:32.963033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_sample(hubmap_dataset, n_sample):\n    n_rows = n_sample\n    n_cols = 3\n    ratio = 2\n    \n    fig, axs = plt.subplots(n_rows, n_cols, figsize=(n_cols*ratio, n_rows*ratio))\n    \n    for i in range(n_sample):\n        # Gather samples\n        image, mask = hubmap_dataset[i]\n    \n        # shape (H, W, C)\n        image = image.permute(1, 2, 0).cpu().numpy() \n        mask = mask.permute(1, 2, 0).cpu().numpy()\n\n        # Plot\n        if n_sample == 1:\n            axs[0].imshow(image, cmap='gray'); axs[0].axis('off')\n            axs[1].imshow(mask, cmap='gray'); axs[1].axis('off')\n            axs[2].imshow(image, cmap='gray', interpolation=None); axs[2].axis('off')\n            axs[2].imshow(mask, cmap='rainbow', alpha=0.6, interpolation=None)\n            \n            # Set subplot title\n            axs[0].set_title('Image')\n            axs[1].set_title('Mask')\n            axs[2].set_title('Overlay')\n            \n            for ax in axs:\n                # Hide X and Y axes label marks\n                ax.xaxis.set_tick_params(labelbottom=False)\n                ax.yaxis.set_tick_params(labelleft=False)\n\n                # Hide X and Y axes tick marks\n                ax.set_xticks([])\n                ax.set_yticks([])\n        else:\n            axs[i, 0].imshow(image, cmap='gray'); axs[i, 0].axis('off')\n            axs[i, 1].imshow(mask, cmap='gray'); axs[i, 1].axis('off')\n            axs[i, 2].imshow(image, cmap='gray'); axs[i, 2].axis('off')\n            axs[i, 2].imshow(mask, cmap='rainbow', alpha=0.6)\n            \n            # Set column title\n            cols_label = ['Image', 'Mask', 'Overlay']\n            pad = 10\n            \n            for ax, col_label in zip(axs[0], cols_label):\n                ax.annotate(col_label, xy=(0.5, 1), xytext=(0, pad),\n                            xycoords='axes fraction', textcoords='offset points', \n                            ha='center', va='baseline')\n            \n            for ax in axs[i, :]:\n                # Hide X and Y axes label marks\n                ax.xaxis.set_tick_params(labelbottom=False)\n                ax.yaxis.set_tick_params(labelleft=False)\n\n                # Hide X and Y axes tick marks\n                ax.set_xticks([])\n                ax.set_yticks([])","metadata":{"_uuid":"a977c7cd-8893-4a12-9096-3f476f43b07d","_cell_guid":"2a2e3f55-7b6c-491c-90cc-8f917dda176d","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:32.965365Z","iopub.execute_input":"2023-06-19T00:21:32.965738Z","iopub.status.idle":"2023-06-19T00:21:32.981004Z","shell.execute_reply.started":"2023-06-19T00:21:32.965704Z","shell.execute_reply":"2023-06-19T00:21:32.980094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_sample(dataset, 4)","metadata":{"_uuid":"550f58cd-1c0d-496c-8299-039a52bc9d3f","_cell_guid":"ddff27b8-6665-4a98-9ff0-a54bb749f508","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:32.982434Z","iopub.execute_input":"2023-06-19T00:21:32.983031Z","iopub.status.idle":"2023-06-19T00:21:34.368566Z","shell.execute_reply.started":"2023-06-19T00:21:32.982999Z","shell.execute_reply":"2023-06-19T00:21:34.367702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's split our current dataset into training data and validation data","metadata":{"_uuid":"05b3a550-380d-4b20-9205-eef421a3c07e","_cell_guid":"3e129f17-b7c5-4b18-acc8-7275d57d3a51","trusted":true}},{"cell_type":"code","source":"train_data, valid_data = random_split(dataset, [0.8, 0.2])","metadata":{"_uuid":"7047db38-e60d-493e-8d32-60efc19067eb","_cell_guid":"26a283de-5c25-4683-9c85-1eca57012a75","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:34.370123Z","iopub.execute_input":"2023-06-19T00:21:34.370757Z","iopub.status.idle":"2023-06-19T00:21:34.375301Z","shell.execute_reply.started":"2023-06-19T00:21:34.370724Z","shell.execute_reply":"2023-06-19T00:21:34.374492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'No. of images in train_data: {len(train_data)}')\nprint(f'No. of images in val_data: {len(valid_data)}')","metadata":{"_uuid":"dc5a9088-b72a-4f4c-9110-c657bccb9d8d","_cell_guid":"3797ad39-6628-480a-b2dc-e51b21ac347d","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:34.376736Z","iopub.execute_input":"2023-06-19T00:21:34.377392Z","iopub.status.idle":"2023-06-19T00:21:34.391498Z","shell.execute_reply.started":"2023-06-19T00:21:34.377364Z","shell.execute_reply":"2023-06-19T00:21:34.390538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now, it's turn to create our DataLoader.\nDataLoader is an iterable that abstracts the complexity of how we input the data to the model.","metadata":{"_uuid":"1a68f58b-2490-47c6-94a9-0421bca99368","_cell_guid":"2c8f0e4e-13c2-4c72-8a04-f08281fbeedf","trusted":true}},{"cell_type":"code","source":"train_loader = DataLoader(train_data, batch_size=BATCH_SIZE, shuffle=True)\nvalid_loader = DataLoader(valid_data, batch_size=BATCH_SIZE, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.393401Z","iopub.execute_input":"2023-06-19T00:21:34.394177Z","iopub.status.idle":"2023-06-19T00:21:34.399651Z","shell.execute_reply.started":"2023-06-19T00:21:34.394122Z","shell.execute_reply":"2023-06-19T00:21:34.398987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'Total no. of batches in trainloader: {len(train_loader)}')\nprint(f'Total no. of batches in validloader: {len(valid_loader)}')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.401081Z","iopub.execute_input":"2023-06-19T00:21:34.401719Z","iopub.status.idle":"2023-06-19T00:21:34.411374Z","shell.execute_reply.started":"2023-06-19T00:21:34.401688Z","shell.execute_reply":"2023-06-19T00:21:34.410483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for image, mask in train_loader:\n    break\n\nprint(f'One batch image shape: {image.shape}')\nprint(f'One batch mask shape: {mask.shape}')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.412725Z","iopub.execute_input":"2023-06-19T00:21:34.413234Z","iopub.status.idle":"2023-06-19T00:21:34.583750Z","shell.execute_reply.started":"2023-06-19T00:21:34.413198Z","shell.execute_reply":"2023-06-19T00:21:34.581902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Segmentation Model\n","metadata":{"_uuid":"f2e215dc-9b6b-4770-ac23-15d45daf3ad1","_cell_guid":"20e6ad8c-e8b9-4864-ade0-0c8a11233bde","trusted":true}},{"cell_type":"code","source":"class SegmentationModel(nn.Module):\n    def __init__(self):\n        super(SegmentationModel, self).__init__()\n        self.arch = smp.Unet(\n            encoder_name=ENCODER, \n            encoder_weights=WEIGHTS,\n            in_channels=3,\n            classes=1,\n            activation=None\n        )\n\n    def forward(self, images, masks=None):\n        logits = self.arch(images)\n\n        if masks is not None:\n            dice = DiceLoss(mode='binary')(logits, masks)\n            bce = nn.BCEWithLogitsLoss()(logits, masks)\n\n            return logits, dice + bce\n\n        return logits","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.585358Z","iopub.execute_input":"2023-06-19T00:21:34.585702Z","iopub.status.idle":"2023-06-19T00:21:34.593408Z","shell.execute_reply.started":"2023-06-19T00:21:34.585668Z","shell.execute_reply":"2023-06-19T00:21:34.592147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = SegmentationModel()\n\nif torch.cuda.device_count() > 1:\n    print(\"Let's use\", torch.cuda.device_count(), \"GPUs!\")\n    model = nn.DataParallel(model)\n\nmodel.to(DEVICE); # \";\" avoids printing model's internal structure","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.595435Z","iopub.execute_input":"2023-06-19T00:21:34.595890Z","iopub.status.idle":"2023-06-19T00:21:34.806076Z","shell.execute_reply.started":"2023-06-19T00:21:34.595837Z","shell.execute_reply":"2023-06-19T00:21:34.805062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fn(data_loader, model, optimizer):\n    model.train()\n    total_loss = 0.0\n\n    for images, masks in tqdm(data_loader):\n        torch.cuda.empty_cache()\n        images = images.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        optimizer.zero_grad()\n        logits, loss = model(images, masks)\n        \n        if isinstance(model, nn.DataParallel):\n            loss.mean().backward()\n            total_loss += loss.mean().item()\n        else:\n            loss.backward()\n            total_loss += loss.item()\n            \n        optimizer.step()\n        \n    return total_loss / len(data_loader)","metadata":{"_uuid":"923da8e3-ecef-4c78-bea9-3364366f5acc","_cell_guid":"1bb07e9b-9cb5-44b3-b6cb-0acd2e1554c6","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:21:34.807641Z","iopub.execute_input":"2023-06-19T00:21:34.808412Z","iopub.status.idle":"2023-06-19T00:21:34.816159Z","shell.execute_reply.started":"2023-06-19T00:21:34.808377Z","shell.execute_reply":"2023-06-19T00:21:34.815276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eval_fn(data_loader, model):\n    model.eval()\n    total_loss = 0.0\n\n    with torch.no_grad():\n        for images, masks in tqdm(data_loader):\n            torch.cuda.empty_cache()\n            images = images.to(DEVICE)\n            masks = masks.to(DEVICE)\n\n            logits, loss = model(images, masks)\n\n            if isinstance(model, nn.DataParallel):\n                total_loss += loss.mean().item()\n            else:\n                total_loss += loss.item()\n\n    return total_loss / len(data_loader)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.817467Z","iopub.execute_input":"2023-06-19T00:21:34.817871Z","iopub.status.idle":"2023-06-19T00:21:34.828886Z","shell.execute_reply.started":"2023-06-19T00:21:34.817796Z","shell.execute_reply":"2023-06-19T00:21:34.828011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"optimizer = torch.optim.Adam(model.parameters(), lr=LR)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.830419Z","iopub.execute_input":"2023-06-19T00:21:34.830809Z","iopub.status.idle":"2023-06-19T00:21:34.847489Z","shell.execute_reply.started":"2023-06-19T00:21:34.830778Z","shell.execute_reply":"2023-06-19T00:21:34.846493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_valid_loss = np.Inf\ntrain_losses = []\nvalid_losses = []\n\nfor i in range(EPOCHS):\n    train_loss = train_fn(train_loader, model, optimizer)\n    valid_loss = eval_fn(valid_loader, model)\n\n    train_losses.append(train_loss)\n    valid_losses.append(valid_loss)\n\n    if valid_loss < best_valid_loss:\n        torch.save(model.state_dict(), MODEL_FILE)\n        print('SAVED MODEL')\n        best_valid_loss = valid_loss\n\n    print(f'Epoch: {i+1} Training Loss: {train_loss} Validation Loss: {valid_loss}')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:21:34.848698Z","iopub.execute_input":"2023-06-19T00:21:34.849098Z","iopub.status.idle":"2023-06-19T00:47:18.665980Z","shell.execute_reply.started":"2023-06-19T00:21:34.849064Z","shell.execute_reply":"2023-06-19T00:47:18.664915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Generate a sequence of integers to represent the epoch numbers\nepochs = range(1, EPOCHS+1)\n \n# Plot and label the training and validation loss values\nplt.plot(epochs, train_losses, label='Training Loss')\nplt.plot(epochs, valid_losses, label='Validation Loss')\n \n# Add in a title and axes labels\nplt.title('Training and Validation Loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\n \n# Set the tick locations\nplt.xticks(np.arange(0, EPOCHS+1, 1))\n \n# Display the plot\nplt.legend(loc='best')\nplt.show()","metadata":{"_uuid":"28591bd2-8e79-4f21-bd00-8991cbbe5543","_cell_guid":"e587fa51-9822-4279-8342-55e15657d65f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:47:18.667485Z","iopub.execute_input":"2023-06-19T00:47:18.669070Z","iopub.status.idle":"2023-06-19T00:47:19.043709Z","shell.execute_reply.started":"2023-06-19T00:47:18.669034Z","shell.execute_reply":"2023-06-19T00:47:19.042731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load(MODEL_FILE))","metadata":{"_uuid":"6c60e3d7-9c1c-44a9-927a-c3b4a55372ad","_cell_guid":"64ce73df-9c45-48c6-840a-62dd922fa89f","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:47:19.049027Z","iopub.execute_input":"2023-06-19T00:47:19.049313Z","iopub.status.idle":"2023-06-19T00:47:19.146225Z","shell.execute_reply.started":"2023-06-19T00:47:19.049282Z","shell.execute_reply":"2023-06-19T00:47:19.144794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx = 8\n\nimage, mask = valid_data[idx]\nlogits_mask = model(image.to(DEVICE).unsqueeze(0)) # (C, H, W) -> (1, C, H, W)\npred_mask = torch.sigmoid(logits_mask)\npred_mask = pred_mask.squeeze(0)","metadata":{"_uuid":"b041e7af-1a70-4b09-81a4-961aa95be0c9","_cell_guid":"971f05da-7e6e-4b04-a9fe-ce5bf1e3220e","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:47:19.148104Z","iopub.execute_input":"2023-06-19T00:47:19.148465Z","iopub.status.idle":"2023-06-19T00:47:19.356393Z","shell.execute_reply.started":"2023-06-19T00:47:19.148432Z","shell.execute_reply":"2023-06-19T00:47:19.355365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'img shape {image.shape}')\nprint(f'mask shape {mask.shape}')\nprint(f'pred_mask shape {pred_mask.shape}')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:47:19.358154Z","iopub.execute_input":"2023-06-19T00:47:19.358510Z","iopub.status.idle":"2023-06-19T00:47:19.365208Z","shell.execute_reply.started":"2023-06-19T00:47:19.358479Z","shell.execute_reply":"2023-06-19T00:47:19.363797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot inference \nn_rows = 1\nn_cols = 3\nratio = 2\n\nfig, axs = plt.subplots(n_rows, n_cols, figsize=(n_cols*ratio, n_rows*ratio))\n\n# shape (H, W, C)\nimage = image.permute(1, 2, 0).cpu().numpy() \nmask = mask.permute(1, 2, 0).cpu().numpy()\npred_mask = pred_mask.permute(1, 2, 0).detach().cpu().numpy()\n\n# Plot\naxs[0].imshow(image); axs[0].axis('off')\naxs[1].imshow(mask, cmap='plasma'); axs[1].axis('off')\naxs[2].imshow(pred_mask, cmap='plasma', interpolation=None); axs[2].axis('off')\n\n# Set subplot title\naxs[0].set_title('Image')\naxs[1].set_title('Ground Truth')\naxs[2].set_title('Inference')\n\nfor ax in axs:\n    # Hide X and Y axes label marks\n    ax.xaxis.set_tick_params(labelbottom=False)\n    ax.yaxis.set_tick_params(labelleft=False)\n\n    # Hide X and Y axes tick marks\n    ax.set_xticks([])\n    ax.set_yticks([])","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:47:19.367247Z","iopub.execute_input":"2023-06-19T00:47:19.368702Z","iopub.status.idle":"2023-06-19T00:47:19.715617Z","shell.execute_reply.started":"2023-06-19T00:47:19.368637Z","shell.execute_reply":"2023-06-19T00:47:19.714743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = '/kaggle/input/hubmap-hacking-the-human-vasculature/test/72e40acccadf.tif'\n        \nwith rasterio.open(img_path) as image:\n    image_array = image.read()","metadata":{"_uuid":"fe5d51e1-c605-45ab-a627-3f6d2f1c18c1","_cell_guid":"e9603bbd-e5aa-40fb-8cab-912a4f6742a3","collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2023-06-19T00:47:19.717394Z","iopub.execute_input":"2023-06-19T00:47:19.718210Z","iopub.status.idle":"2023-06-19T00:47:19.754415Z","shell.execute_reply.started":"2023-06-19T00:47:19.718176Z","shell.execute_reply":"2023-06-19T00:47:19.753544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pipreqsnb '/kaggle/working'","metadata":{"execution":{"iopub.status.busy":"2023-06-19T00:47:19.755939Z","iopub.execute_input":"2023-06-19T00:47:19.756285Z","iopub.status.idle":"2023-06-19T00:47:21.171675Z","shell.execute_reply.started":"2023-06-19T00:47:19.756253Z","shell.execute_reply":"2023-06-19T00:47:21.170514Z"},"trusted":true},"execution_count":null,"outputs":[]}]}