{"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":"# Efficient PyTorch Dataloader for HubMap Competition\n\nHello Kagglers! In this notebook, we introduce a custom PyTorch Dataloader for the HubMap: Hacking the  Hacking the Human Vasculature competition. The main challenge here is handling high-resolution TIFF images alongside complex polygonal segmentation masks. \n\nOur Dataloader manages this complexity for you, efficiently loading image tiles and corresponding segmentation masks in an appropriate format for direct input into your PyTorch models. **The dataloader only loads the 'blood_sample' mask, but the same logic can be extended if you want the other masks**.\n\nCrucially, it outputs tensors of the shape `[BS, C, H, W]` for images and `[BS, H, W]` for masks where `BS` stands for batch size, `C` stands for channels, `H` for height, and `W` for width.\n\nFeel free to incorporate this Dataloader into your workflow, and good luck with the competition!\n\n--- \n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-23T20:08:51.991201Z","iopub.execute_input":"2023-05-23T20:08:51.991603Z","iopub.status.idle":"2023-05-23T20:08:52.005606Z","shell.execute_reply.started":"2023-05-23T20:08:51.991574Z","shell.execute_reply":"2023-05-23T20:08:52.004451Z"}}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\nfrom PIL import Image\nimport os\nimport json\nimport numpy as np\nfrom skimage.draw import polygon2mask\n\nclass CustomDataset(Dataset):\n    def __init__(self, image_dir, labels_file, transform=None):\n        with open(labels_file, 'r') as json_file:\n            self.json_labels = [json.loads(line) for line in json_file]\n\n        self.image_dir = image_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.json_labels)\n\n    def __getitem__(self, idx):\n        # Load image\n        image_path = os.path.join(self.image_dir, f\"{self.json_labels[idx]['id']}.tif\")\n        image = Image.open(image_path)\n\n        # Initialize mask\n        mask = np.zeros((512, 512), dtype=np.float32)\n\n        # Process annotations\n        for annot in self.json_labels[idx]['annotations']:\n            cords = annot['coordinates']\n            if annot['type'] == \"blood_vessel\":\n                for cord in cords:\n                    rr, cc = np.array([i[1] for i in cord]), np.asarray([i[0] for i in cord])\n                    mask[rr, cc] = 1\n\n        # Convert PIL Image and mask to PyTorch tensor\n        image = torch.tensor(np.array(image), dtype=torch.float32).permute(2, 0, 1)  # Shape: [C, H, W]\n        mask = torch.tensor(mask, dtype=torch.float32)\n\n        if self.transform:\n            image = self.transform(image)\n            mask = self.transform(mask)\n        return image, mask","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-06T19:51:59.947115Z","iopub.execute_input":"2023-06-06T19:51:59.947783Z","iopub.status.idle":"2023-06-06T19:52:04.683517Z","shell.execute_reply.started":"2023-06-06T19:51:59.947750Z","shell.execute_reply":"2023-06-06T19:52:04.682339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n# Instantiate the dataset\ndataset = CustomDataset(image_dir='../input/hubmap-hacking-the-human-vasculature/train/', \n                        labels_file='../input/hubmap-hacking-the-human-vasculature/polygons.jsonl')\n\n# Instantiate the DataLoader and load a sample batch for viz:\ndataloader = DataLoader(dataset, batch_size=32, shuffle=True)\nimage, mask = next(iter(dataloader))\nprint(image.shape)\nprint(mask.shape)","metadata":{"execution":{"iopub.status.busy":"2023-06-06T19:52:04.685448Z","iopub.execute_input":"2023-06-06T19:52:04.686079Z","iopub.status.idle":"2023-06-06T19:52:11.098300Z","shell.execute_reply.started":"2023-06-06T19:52:04.686047Z","shell.execute_reply":"2023-06-06T19:52:11.097106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef plot_image_and_mask(image, mask):\n    fig, ax = plt.subplots(1, 2, figsize=(12, 6))\n\n    # Transpose the dimensions to height, width, channels for plotting\n    image = image.transpose((1, 2, 0))\n\n    # Plot the image\n    ax[0].imshow(image)  \n    ax[0].set_title(\"Image\")\n    ax[0].axis('off')\n\n    # Plot the mask\n    ax[1].imshow(mask, cmap='gray')\n    ax[1].set_title(\"Mask 'Blood_vessel'\")\n    ax[1].axis('off')\n\n    plt.show()\n\nplot_image_and_mask(image[10].numpy()/255, mask[10].numpy())\n","metadata":{"execution":{"iopub.status.busy":"2023-06-06T19:53:30.128907Z","iopub.execute_input":"2023-06-06T19:53:30.129320Z","iopub.status.idle":"2023-06-06T19:53:30.557068Z","shell.execute_reply.started":"2023-06-06T19:53:30.129289Z","shell.execute_reply":"2023-06-06T19:53:30.555771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}