{"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":"# Information\n\n**This notebook provided a basic baseline training with UNET model from segmentation models pytorch. I trained this model with the parameters as set in this notebook and scores with 0.624 in the public leaderboard. I will also provide a notebook for submitting the trained model!**","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport os\nimport numpy as np\n!pip install -q segmentation_models_pytorch\nimport segmentation_models_pytorch as smp\nfrom tqdm.notebook import tqdm\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nimport torch.utils.checkpoint as C\nimport torchvision.transforms.functional as fn\nimport torchvision.transforms as T\nimport matplotlib.pyplot as plt\n!pip install -q torchsummary\nfrom torchvision import models\nfrom torchsummary import summary","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get the Device","metadata":{}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = torch.device('cuda')\nelse:\n    device = torch.device('cpu')\n    \ndevice","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:12.59138Z","iopub.execute_input":"2023-06-02T14:25:12.591761Z","iopub.status.idle":"2023-06-02T14:25:12.624193Z","shell.execute_reply.started":"2023-06-02T14:25:12.591722Z","shell.execute_reply":"2023-06-02T14:25:12.622969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config File","metadata":{}},{"cell_type":"code","source":"class CFG:\n    \n    # Path to the data folder (Thanks to @Kenni)\n    GLOBAL_PATH = '/kaggle/input/google-research-identify-contrails-preprocessing'\n    \n    # base image size\n    resize_value = 256\n    \n    # resize image\n    resize = True\n    if resize:\n        resize_value = 384\n        \n    # Model Settings    \n    model = 'UNET'\n    encoder = 'efficientnet-b0'\n    weights = 'imagenet'\n    \n    batch_size = 16\n    optimizer='Adam'\n    lr = 5e-4\n    epochs = 20","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:12.626629Z","iopub.execute_input":"2023-06-02T14:25:12.626905Z","iopub.status.idle":"2023-06-02T14:25:12.634021Z","shell.execute_reply.started":"2023-06-02T14:25:12.626882Z","shell.execute_reply":"2023-06-02T14:25:12.632446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the Torch Dataset","metadata":{}},{"cell_type":"code","source":"#A custom Dataset class must implement three functions: __init__, __len__, and __getitem__\nclass ContrailDataset(Dataset):\n    \n    def __init__(self, base_dir, data_type='train'):\n        assert data_type in ['train_images', 'validate_images'], \\\n            \"'data_type' should be one of 'train_images' or 'validate_images'\"\n        \n        self.base_dir = base_dir\n        self.data_type = data_type\n        self.record = os.listdir(self.base_dir +'/'+ self.data_type)\n       \n        self.resize_image = T.Resize(CFG.resize_value,interpolation=T.InterpolationMode.BILINEAR,antialias=True)\n        self.resize_mask = T.Resize(CFG.resize_value,interpolation=T.InterpolationMode.NEAREST,antialias=True)\n   \n    def __len__(self):\n        return len(self.record)\n\n    def __getitem__(self, idx):\n        \n        record_id = self.record[idx]\n        record_dir = os.path.join(self.base_dir, self.data_type, record_id)\n        \n        false_color = np.load(os.path.join(record_dir,'image.npy'))\n        human_pixel_mask = np.load(os.path.join(record_dir,'human_pixel_masks.npy')) \n        \n        false_color = torch.from_numpy(false_color)#.clone().detach()\n        human_pixel_mask = torch.from_numpy(human_pixel_mask)#.clone().detach()\n        \n        false_color = torch.moveaxis(false_color,-1,0)\n        human_pixel_mask = torch.moveaxis(human_pixel_mask,-1,0)\n            \n        if self.data_type == 'train':\n            \n            random_crop_factor = torch.rand(1)\n            crop_min, crop_max = 0.5 , 1\n            crop_factor = crop_min + random_crop_factor * (crop_max-crop_min) \n            crop_size = int(crop_factor * 256)\n            self.crop = T.CenterCrop(size=crop_size)\n            \n            false_color = self.crop(false_color)\n            human_pixel_mask =  self.crop(human_pixel_mask)\n            \n            false_color = self.resize_image(false_color)\n            human_pixel_mask =  self.resize_mask(human_pixel_mask)\n\n        \n        #if CFG.resize and self.data_type=='validation':\n            #false_color = self.resize_image(false_color)\n            #human_pixel_mask =  self.resize_mask(human_pixel_mask)\n                  \n        # false color is scaled between 0 and 1!\n        return false_color, human_pixel_mask.float()\n","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:12.649149Z","iopub.execute_input":"2023-06-02T14:25:12.649541Z","iopub.status.idle":"2023-06-02T14:25:12.665033Z","shell.execute_reply.started":"2023-06-02T14:25:12.649511Z","shell.execute_reply":"2023-06-02T14:25:12.663894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the Training and Validation Dataloader","metadata":{}},{"cell_type":"code","source":"training_data = ContrailDataset(base_dir=CFG.GLOBAL_PATH, data_type='train_images')\ntrain_dataloader = DataLoader(\n    training_data, \n    batch_size=CFG.batch_size, \n    shuffle=True, \n    num_workers= 4 if torch.cuda.is_available() else 0,\n    pin_memory=True,\n    drop_last = True\n)\n\nvalidation_data = ContrailDataset(base_dir=CFG.GLOBAL_PATH, data_type='validate_images')\nvalidation_dataloader = DataLoader(\n    validation_data, \n    batch_size=CFG.batch_size, \n    shuffle=False, \n    num_workers= 4 if torch.cuda.is_available() else 0,\n    pin_memory=True,\n    drop_last = True\n)","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:12.666596Z","iopub.execute_input":"2023-06-02T14:25:12.666953Z","iopub.status.idle":"2023-06-02T14:25:13.091731Z","shell.execute_reply.started":"2023-06-02T14:25:12.666921Z","shell.execute_reply":"2023-06-02T14:25:13.08772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show some Images from the Dataloaders","metadata":{}},{"cell_type":"code","source":"image,mask = next(iter(validation_dataloader))\n\nimage = torch.moveaxis(image,1,-1)\nmask = torch.moveaxis(mask,1,-1)\n\nfor i in range(1):\n\n    plt.figure(figsize=(18, 6))\n    \n    ax = plt.subplot(1, 3, 1)\n    ax.imshow(image[i])\n    ax.set_title('False color image')\n    \n\n    ax = plt.subplot(1, 3, 2)\n    ax.imshow(mask[i], interpolation='none')\n    ax.set_title('Ground truth contrail mask')\n        \n    ax = plt.subplot(1, 3, 3)\n    ax.imshow(image[i])\n    ax.imshow(mask[i], cmap='Reds', alpha=.4, interpolation='none')\n    ax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:13.095736Z","iopub.execute_input":"2023-06-02T14:25:13.096497Z","iopub.status.idle":"2023-06-02T14:25:24.714906Z","shell.execute_reply.started":"2023-06-02T14:25:13.096457Z","shell.execute_reply":"2023-06-02T14:25:24.711022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the Model (from SMP)","metadata":{}},{"cell_type":"code","source":"if CFG.model == 'UNET':\n    model = smp.Unet(\n    encoder_name =CFG.encoder,\n    encoder_weights=CFG.weights,    # use `imagenet` pre-trained weights for encoder initialization\n    in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n    classes=1,        # model output channels (number of classes in your dataset)\n    activation=\"sigmoid\",\n    )\n    model.to(device)\n    summary(model, (3, 256, 256))\n","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:33.516434Z","iopub.execute_input":"2023-06-02T14:25:33.517632Z","iopub.status.idle":"2023-06-02T14:25:39.273518Z","shell.execute_reply.started":"2023-06-02T14:25:33.517588Z","shell.execute_reply":"2023-06-02T14:25:39.272596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimizer","metadata":{}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=CFG.lr)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min',patience=2)\nprint(f'learning rate: {optimizer.param_groups[0][\"lr\"]}')","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:39.283607Z","iopub.execute_input":"2023-06-02T14:25:39.284095Z","iopub.status.idle":"2023-06-02T14:25:39.304176Z","shell.execute_reply.started":"2023-06-02T14:25:39.284062Z","shell.execute_reply":"2023-06-02T14:25:39.303135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss Function","metadata":{}},{"cell_type":"code","source":"def dice_global(y_p,y_t,smooth=1e-3):\n\n    intersection = torch.sum(y_p * y_t)\n    union = torch.sum(y_p) + torch.sum(y_t)\n\n    dice = (2.0 * intersection + smooth) / (union + smooth)\n\n    return dice\n\ndef dice_loss_global(y_p,y_t):\n    return 1-dice_global(y_p,y_t)","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:39.31596Z","iopub.execute_input":"2023-06-02T14:25:39.316384Z","iopub.status.idle":"2023-06-02T14:25:39.325961Z","shell.execute_reply.started":"2023-06-02T14:25:39.316347Z","shell.execute_reply":"2023-06-02T14:25:39.324816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Average dice score for the examples in a batch\ndef dice_avg(y_p, y_t,smooth=1e-3):\n    i = torch.sum(y_p * y_t, dim=(2, 3))\n    u = torch.sum(y_p, dim=(2, 3)) + torch.sum(y_t, dim=(2, 3))\n    score = (2 * i + smooth)/(u + smooth)\n    return torch.mean(score)\n\n\ndef dice_loss_avg(y_p,y_t):\n    return 1-dice_score_jan(y_p,y_t)","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:39.327483Z","iopub.execute_input":"2023-06-02T14:25:39.328046Z","iopub.status.idle":"2023-06-02T14:25:39.34105Z","shell.execute_reply.started":"2023-06-02T14:25:39.327941Z","shell.execute_reply":"2023-06-02T14:25:39.339982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training and Validation Loop","metadata":{}},{"cell_type":"code","source":"train_dice_global = []\ntrain_dice_avg = []\neval_dice_global = []\neval_dice_avg = []\nbst_dice = 0\nbst_epoch = 1\nfor epoch in range(1,CFG.epochs+1):\n    \n    print(f'________epoch: {epoch}________')\n    \n    # Early stopping\n    if epoch-bst_epoch >=5:\n        print(f'early stopping in epoch {epoch}')\n        break\n    \n    model.train()\n    bar = tqdm(train_dataloader)\n    tot_loss_global = 0\n    tot_dice_global = 0\n    tot_dice_avg = 0\n    count = 0\n    for image, mask in bar:\n        \n        image = torch.nn.functional.interpolate(image, \n                                                size=CFG.resize_value,\n                                                mode='bilinear'\n                                               )\n        \n        # Transfer to Device\n        image,mask = image.to(device), mask.to(device)\n        \n        # Set optimizer gradients to zero\n        optimizer.zero_grad()\n        \n        #Perform Inference\n        pred_mask = model(image)\n        \n        # If the image was resized, use a resizing step to make 256 again\n        if CFG.resize:\n            pred_mask = torch.nn.functional.interpolate(pred_mask, \n                                                        size=256,\n                                                        mode='bilinear'\n                                                       )\n        \n        # Calculate the loss and do a backward pass\n        loss = dice_loss_global(pred_mask, mask)\n        loss.backward()\n        \n        # Adjust the weights\n        optimizer.step()\n\n        tot_loss_global += loss.item()\n        tot_dice_global+=1-loss.item()\n        tot_dice_avg += dice_avg(pred_mask,mask).item()\n        count += 1\n        bar.set_postfix(TrainDiceLossGlobal=f'{tot_loss_global/count:.4f}', \n                        TrainDiceGlobal=f'{tot_dice_global/count:.4f}',\n                        TrainDiceAvg = f'{tot_dice_avg/count:.4f}')\n        \n    train_dice_global.append(np.array(tot_dice_global/count))\n    train_dice_avg.append(np.array(tot_dice_avg/count))\n      \n    model.train(False)\n    bar = tqdm(validation_dataloader)\n    tot_dice_global = 0\n    tot_dice_avg = 0\n    count = 0\n    for image, mask in bar:\n        \n        if CFG.resize:\n            image = torch.nn.functional.interpolate(image, \n                                                size=CFG.resize_value,\n                                                mode='bilinear'\n                                               )\n        image,mask = image.to(device), mask.to(device)\n        pred_mask = model(image)\n        \n        if CFG.resize:\n            pred_mask = torch.nn.functional.interpolate(pred_mask, \n                                                size=256,\n                                                mode='bilinear'\n                                               )\n        \n        tot_dice_global += dice_global(pred_mask, mask).item()\n        tot_dice_avg+=dice_avg(pred_mask,mask).item()\n        count += 1\n        bar.set_postfix(ValidDiceGlobal=f'{tot_dice_global/count:.4f}',\n                        ValidDiceAvg = f'{tot_dice_avg/count:.4f}')\n        \n\n    eval_dice_global.append(np.array(tot_dice_global/count))\n    eval_dice_avg.append(np.array(tot_dice_avg/count))\n    scheduler.step(1-(tot_dice_global/count))\n    print(f'learning rate: {optimizer.param_groups[0][\"lr\"]}')\n        \n    if tot_dice_global/count > bst_dice:\n        bst_dice = tot_dice_global/count\n        bst_epoch = epoch\n        torch.save(model.state_dict(), f'model_state_dict_epoch_{epoch}_dice_{bst_dice:.4f}.pth')\n        torch.save(model, f'model_epoch_{epoch}_dice_{bst_dice:.4f}.pt')\n        print(f\"current model saved! Epoch: {epoch} global dice: {bst_dice} avg dice: {tot_dice_avg/count}\") \n        \n ","metadata":{"execution":{"iopub.status.busy":"2023-06-02T14:25:39.366433Z","iopub.execute_input":"2023-06-02T14:25:39.367252Z","iopub.status.idle":"2023-06-02T15:49:27.890915Z","shell.execute_reply.started":"2023-06-02T14:25:39.367222Z","shell.execute_reply":"2023-06-02T15:49:27.887877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training and Validation History","metadata":{}},{"cell_type":"code","source":"plt.plot(train_dice_global, label='train_dice_global')\nplt.plot(train_dice_avg,label='train_dice_avg')\nplt.plot(eval_dice_global, label='eval_dice_global')\nplt.plot(eval_dice_avg,label='eval_dice_avg')\nplt.legend()\nplt.show","metadata":{"execution":{"iopub.status.busy":"2023-06-02T15:49:27.895244Z","iopub.execute_input":"2023-06-02T15:49:27.895758Z","iopub.status.idle":"2023-06-02T15:49:28.360771Z","shell.execute_reply.started":"2023-06-02T15:49:27.895696Z","shell.execute_reply":"2023-06-02T15:49:28.358129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show some predictions for the validation dataset","metadata":{}},{"cell_type":"code","source":"image,mask = next(iter(validation_dataloader))\n\nimage,mask = image.to(device), mask.to(device)\npred_mask = model(image)\n\nimage = torch.moveaxis(image,1,-1)\nmask = torch.moveaxis(mask,1,-1)\npred_mask = torch.moveaxis(pred_mask,1,-1)\n\nimage, mask, pred_mask = image.cpu(),mask.cpu(),pred_mask.detach().cpu()\n\nfor i in range(CFG.batch_size):\n    \n    plt.figure(figsize=(18, 6))\n    \n    ax = plt.subplot(1, 3, 1)\n    ax.imshow(image[i])\n    ax.set_title('False color image')\n    \n\n    ax = plt.subplot(1, 3, 2)\n    ax.imshow(mask[i], interpolation='none')\n    ax.set_title('Ground truth contrail mask')\n    \n    ax = plt.subplot(1, 3, 3)\n    ax.imshow(pred_mask[i], interpolation='none')\n    ax.set_title('Predicted_Mask')\n        \n","metadata":{"execution":{"iopub.status.busy":"2023-06-02T15:49:28.369463Z","iopub.execute_input":"2023-06-02T15:49:28.370649Z","iopub.status.idle":"2023-06-02T15:50:01.921818Z","shell.execute_reply.started":"2023-06-02T15:49:28.37061Z","shell.execute_reply":"2023-06-02T15:50:01.920602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Bonus: Vary the threshold for the predictions","metadata":{}},{"cell_type":"code","source":"bst_dice = 0\nthresholds = [0.01,0.02,0.05,0.1,0.2,0.3,0.4,0.5,0.6,0.7,0.8,0.9]\nfor threshold in thresholds:\n    \n    model.train(False)\n    bar = tqdm(validation_dataloader)\n\n    tot_dice_avg = 0\n    tot_dice_global = 0\n    count = 0\n    for image, mask in bar:\n        \n        if CFG.resize:\n            image = torch.nn.functional.interpolate(image, \n                                                size=CFG.resize_value,\n                                                mode='bilinear'\n                                               )\n        image,mask = image.to(device), mask.to(device)\n        pred_mask = model(image)\n        \n        \n        pred_mask[pred_mask >= threshold] = 1\n        pred_mask[pred_mask<threshold]=0\n        \n        if CFG.resize:\n            pred_mask = torch.nn.functional.interpolate(pred_mask, \n                                                size=256,\n                                                mode='bilinear'\n                                               )\n        \n        tot_dice_avg += dice_avg(pred_mask, mask).item()\n        tot_dice_global+=dice_global(pred_mask,mask).item()\n        count += 1\n        bar.set_postfix(ValidDiceAvg=f'{tot_dice_avg/count:.4f}',\n                        ValidDiceGlobal = f'{tot_dice_global/count:.4f}')\n        \n \n    if tot_dice_global/count > bst_dice:\n        bst_dice = tot_dice_global/count\n        print(f\"new best global dice: {bst_dice} for threshold: {threshold}\") ","metadata":{"execution":{"iopub.status.busy":"2023-06-02T15:50:01.924006Z","iopub.execute_input":"2023-06-02T15:50:01.924708Z","iopub.status.idle":"2023-06-02T15:52:03.889274Z","shell.execute_reply.started":"2023-06-02T15:50:01.924672Z","shell.execute_reply":"2023-06-02T15:52:03.887726Z"},"trusted":true},"execution_count":null,"outputs":[]}]}