{"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":"<div style=\"background-image: url('https://storage.googleapis.com/kaggle-competitions/kaggle/52279/logos/header.png?t=2023-05-04-23-04-20'); height: 200px; width: 100%;\">\n    <h1 style=\"color: Yellow\">HuBMAP - Hacking the Human Vasculature</h1>\n</div>","metadata":{}},{"cell_type":"code","source":"!pip install -qq git+https://github.com/qubvel/segmentation_models.pytorch","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:47:38.458798Z","iopub.execute_input":"2023-05-24T10:47:38.459282Z","iopub.status.idle":"2023-05-24T10:48:11.047716Z","shell.execute_reply.started":"2023-05-24T10:47:38.459236Z","shell.execute_reply":"2023-05-24T10:48:11.046306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Imports</center></h3>","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport glob\nimport json\n\nimport torch\nimport torch.nn as nn\nimport albumentations as A\n\nimport cv2\nimport pandas as pd\nfrom PIL import Image\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nimport tqdm\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import train_test_split\nfrom skimage.draw import polygon2mask\nimport tifffile as tiff\n\ntorch.backends.cudnn.benchmark = True","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-24T10:48:11.050721Z","iopub.execute_input":"2023-05-24T10:48:11.051138Z","iopub.status.idle":"2023-05-24T10:48:18.246061Z","shell.execute_reply.started":"2023-05-24T10:48:11.051094Z","shell.execute_reply":"2023-05-24T10:48:18.245033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Configuration</center></h3>","metadata":{}},{"cell_type":"code","source":"imsize = 224\nBATCH_SIZE = 16\nDEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\nEPOCHS = 5\nNUM_WORKERS = 4\nSEED = 24\nTRAIN_PATH = '/kaggle/input/hubmap-hacking-the-human-vasculature/train/'\nTEST_PATH = \"/kaggle/input/hubmap-hacking-the-human-vasculature/test/\"\nJSON_FILE = \"/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl\"","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:18.247502Z","iopub.execute_input":"2023-05-24T10:48:18.247869Z","iopub.status.idle":"2023-05-24T10:48:18.276542Z","shell.execute_reply.started":"2023-05-24T10:48:18.247835Z","shell.execute_reply":"2023-05-24T10:48:18.275587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=42):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(42)","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:18.280428Z","iopub.execute_input":"2023-05-24T10:48:18.281717Z","iopub.status.idle":"2023-05-24T10:48:18.294847Z","shell.execute_reply.started":"2023-05-24T10:48:18.281672Z","shell.execute_reply":"2023-05-24T10:48:18.293878Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Dataset</center></h3>","metadata":{}},{"cell_type":"code","source":"class HuBMAPDataset(torch.utils.data.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        self.image_size = imsize\n        \n    def __len__(self):\n        return len(self.json_labels)\n    \n    def img2tensor(self, img,dtype:np.dtype=np.float32):\n        if img.ndim==2 : img = np.expand_dims(img,2)\n        img = np.transpose(img,(2,0,1)) # C , H , W\n        return torch.from_numpy(img.astype(dtype, copy=False))\n    \n    def resize(self, img, interp):\n        return  cv2.resize(\n            img, (self.image_size, self.image_size), interpolation=interp)\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 = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\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        if self.transform:\n            augmented = self.transform(image=image, mask=mask)\n            image, mask = augmented['image'],augmented['mask']\n        \n        return self.img2tensor(self.resize(image , cv2.INTER_NEAREST)) , self.img2tensor(self.resize(mask , cv2.INTER_NEAREST))","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:18.296561Z","iopub.execute_input":"2023-05-24T10:48:18.297369Z","iopub.status.idle":"2023-05-24T10:48:18.311985Z","shell.execute_reply.started":"2023-05-24T10:48:18.297330Z","shell.execute_reply":"2023-05-24T10:48:18.311042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Augmentation</center></h3>","metadata":{}},{"cell_type":"code","source":"def get_aug(p=1.0):\n    return A.Compose([\n        A.HorizontalFlip(),\n        A.VerticalFlip(),\n        A.RandomRotate90(),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9,\n                         border_mode=cv2.BORDER_REFLECT),\n        A.OneOf([\n            A.ElasticTransform(p=.3),\n            A.GaussianBlur(p=.3),\n            A.GaussNoise(p=.3),\n            A.OpticalDistortion(p=0.3),\n            A.GridDistortion(p=.1),\n        ], p=0.3),\n        A.OneOf([\n            A.HueSaturationValue(15,25,0),\n            A.CLAHE(clip_limit=2),\n            A.RandomBrightnessContrast(brightness_limit=0.3, contrast_limit=0.3),\n        ], p=0.3),\n    ], p=p)","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:18.313960Z","iopub.execute_input":"2023-05-24T10:48:18.314274Z","iopub.status.idle":"2023-05-24T10:48:18.325209Z","shell.execute_reply.started":"2023-05-24T10:48:18.314239Z","shell.execute_reply":"2023-05-24T10:48:18.324253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = HuBMAPDataset(TRAIN_PATH, JSON_FILE, transform=get_aug())\ndataset = torch.utils.data.DataLoader(ds,batch_size=8,shuffle=False,num_workers=NUM_WORKERS)\nimgs, masks = next(iter(dataset))","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:18.327817Z","iopub.execute_input":"2023-05-24T10:48:18.328279Z","iopub.status.idle":"2023-05-24T10:48:25.252031Z","shell.execute_reply.started":"2023-05-24T10:48:18.328245Z","shell.execute_reply":"2023-05-24T10:48:25.250615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Visualize</center></h3>","metadata":{}},{"cell_type":"code","source":"# Get five samples from the dataloader\nnum_samples = 5\nsample_imgs = imgs[:num_samples]\nsample_masks = masks[:num_samples]\n\n# Plot the images and masks\nfig, axes = plt.subplots(num_samples, 2, figsize=(10, 10))\n\nfor i in range(num_samples):\n    # Plot image\n    axes[i, 0].imshow(((sample_imgs[i].permute(1,2,0))).numpy().astype(np.uint8))\n    axes[i, 0].set_title('Image')\n    axes[i, 0].axis('off')\n    \n    # Plot mask\n    axes[i, 1].imshow(sample_masks[i].squeeze(), cmap='gray')\n    axes[i, 1].set_title('Mask')\n    axes[i, 1].axis('off')\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:25.253880Z","iopub.execute_input":"2023-05-24T10:48:25.254526Z","iopub.status.idle":"2023-05-24T10:48:26.094384Z","shell.execute_reply.started":"2023-05-24T10:48:25.254478Z","shell.execute_reply":"2023-05-24T10:48:26.093560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Model</center></h3>","metadata":{}},{"cell_type":"code","source":"ENCODER = 'efficientnet-b3'\nENCODER_WEIGHTS = 'imagenet'\nACTIVATION = 'sigmoid' \n\nmodel = smp.Unet(encoder_name=ENCODER, encoder_weights=ENCODER_WEIGHTS, activation=ACTIVATION)\nmodel = model.cuda()","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:49:15.825602Z","iopub.execute_input":"2023-05-24T10:49:15.826339Z","iopub.status.idle":"2023-05-24T10:49:16.288778Z","shell.execute_reply.started":"2023-05-24T10:49:15.826288Z","shell.execute_reply":"2023-05-24T10:49:16.287744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Loss Function</center></h3>","metadata":{}},{"cell_type":"code","source":"class CustomLoss(nn.Module):\n    def __init__(self):\n        super(CustomLoss,self).__init__()\n        self.diceloss = smp.losses.DiceLoss(mode='binary')\n        self.binloss = smp.losses.SoftBCEWithLogitsLoss(reduction = 'mean' , smooth_factor = 0.1)\n\n    def forward(self, output, mask):\n        dice = self.diceloss(outputs,mask)\n        bce = self.binloss(outputs , mask)\n        loss = dice * 0.7 + bce * 0.3\n        return loss","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:27.921968Z","iopub.execute_input":"2023-05-24T10:48:27.922382Z","iopub.status.idle":"2023-05-24T10:48:27.929317Z","shell.execute_reply.started":"2023-05-24T10:48:27.922353Z","shell.execute_reply":"2023-05-24T10:48:27.928121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceCoef(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super().__init__()\n\n    def forward(self, y_pred, y_true, smooth=1.):\n        y_true = y_true.view(-1)\n        y_pred = y_pred.view(-1)\n        \n        #Round off y_pred\n        y_pred = torch.round((y_pred - y_pred.min()) / (y_pred.max() - y_pred.min()))\n        \n        intersection = (y_true * y_pred).sum()\n        dice = (2.0*intersection + smooth)/(y_true.sum() + y_pred.sum() + smooth)\n        \n        return dice","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:27.930888Z","iopub.execute_input":"2023-05-24T10:48:27.931299Z","iopub.status.idle":"2023-05-24T10:48:27.944596Z","shell.execute_reply.started":"2023-05-24T10:48:27.931264Z","shell.execute_reply":"2023-05-24T10:48:27.943713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Split the dataset</center></h3>","metadata":{}},{"cell_type":"code","source":"# Define the sizes for training and validation sets\ntrain_size = int(0.8 * len(ds))  # 80% for training\nval_size = len(ds) - train_size  # Remaining 20% for validation\n\ntrain_dataset, val_dataset = torch.utils.data.random_split(ds, [train_size, val_size])\n\n# Create new dataloaders for training and validation sets\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nval_loader = torch.utils.data.DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:48:27.946022Z","iopub.execute_input":"2023-05-24T10:48:27.946435Z","iopub.status.idle":"2023-05-24T10:48:27.957871Z","shell.execute_reply.started":"2023-05-24T10:48:27.946394Z","shell.execute_reply":"2023-05-24T10:48:27.956949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"install\"><center>Training</center></h3>","metadata":{}},{"cell_type":"code","source":"val_losses = []\nlosses = []\ntrain_scores=[]\nval_scores = []\nbest_loss = 999\nbest_score = 0\n    \n        \noptimizer = torch.optim.Adam([\n    {'params': model.decoder.parameters(), 'lr': 5e-5}, \n    {'params': model.encoder.parameters(), 'lr': 8e-5},  \n])\n    \nscheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n                                              max_lr=1e-3, epochs=EPOCHS, steps_per_epoch=len(train_loader))\n    \nloss_func = CustomLoss()\ndice_coe = DiceCoef()","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:49:20.809471Z","iopub.execute_input":"2023-05-24T10:49:20.809869Z","iopub.status.idle":"2023-05-24T10:49:20.821863Z","shell.execute_reply.started":"2023-05-24T10:49:20.809836Z","shell.execute_reply":"2023-05-24T10:49:20.820568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in tqdm.notebook.tqdm(range(EPOCHS)):        \n    model.train()\n    train_loss = 0\n    score = 0\n        \n    for data in tqdm.notebook.tqdm(train_loader ,total = len(train_loader)):\n        optimizer.zero_grad()\n        img, mask = data\n        img = img.to(DEVICE)\n        mask = mask.to(DEVICE)\n        \n        outputs = model(img)  \n\n        loss = loss_func(outputs, mask)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        train_loss += loss.item()\n        score += dice_coe(outputs,mask).item()\n            \n    train_loss /= len(train_loader)\n    score /= len(train_loader)\n    losses.append(train_loss)\n    train_scores.append(score)\n    print(f\"EPOCH: {epoch + 1}, train_loss: {train_loss} , Dice coe : {score} \") #\n        \n        \n    gc.collect()\n    torch.cuda.empty_cache()\n                \n    model.eval()\n        \n    with torch.no_grad():\n            \n        valid_loss = 0\n        val_score = 0\n            \n        for data in val_loader:\n                \n            img, mask = data\n            img = img.to(DEVICE)\n            mask = mask.to(DEVICE)\n\n            outputs = model(img)\n\n            loss = loss_func(outputs, mask)\n            valid_loss += loss.item()\n            val_score += dice_coe(outputs,mask).item()\n                \n        valid_loss /= len(val_loader)\n        val_losses.append(valid_loss)\n            \n        val_score /= len(val_loader)\n        val_scores.append(val_score)\n            \n        print(f\"EPOCH: {epoch + 1}, valid_loss: {valid_loss} , Val Dice COE : {val_score}\") #\n            \n        gc.collect()\n        torch.cuda.empty_cache()\n            \n    if val_score > best_score:\n        best_score = val_score\n        torch.save(model.state_dict(), f\"/kaggle/working/best_score.pth\")\n        print(f\"Saved model for best score : best_score.pth\")\n            \n    if valid_loss < best_loss:\n        best_loss = valid_loss\n        torch.save(model.state_dict(), f\"/kaggle/working/best_loss.pth\")\n        print(f\"Saved model for best loss : best_loss.pth\")    ","metadata":{"execution":{"iopub.status.busy":"2023-05-24T10:57:01.900188Z","iopub.execute_input":"2023-05-24T10:57:01.900551Z","iopub.status.idle":"2023-05-24T10:59:24.967204Z","shell.execute_reply.started":"2023-05-24T10:57:01.900521Z","shell.execute_reply":"2023-05-24T10:59:24.965278Z"},"trusted":true},"execution_count":null,"outputs":[]}]}