{"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* **Reference: [Unet Pytorch Baseline (LB 0.608) - Submission](https://www.kaggle.com/code/janhuebi/unet-pytorch-baseline-lb-0-608-submission)**","metadata":{"papermill":{"duration":0.007327,"end_time":"2023-06-05T23:57:57.008311","exception":false,"start_time":"2023-06-05T23:57:57.000984","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":34.702474,"end_time":"2023-06-05T23:58:31.717638","exception":false,"start_time":"2023-06-05T23:57:57.015164","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:00:47.911929Z","iopub.execute_input":"2023-06-25T07:00:47.912837Z","iopub.status.idle":"2023-06-25T07:01:25.036356Z","shell.execute_reply.started":"2023-06-25T07:00:47.91279Z","shell.execute_reply":"2023-06-25T07:01:25.03514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get the Device","metadata":{"papermill":{"duration":0.006738,"end_time":"2023-06-05T23:58:31.731484","exception":false,"start_time":"2023-06-05T23:58:31.724746","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = torch.device('cuda')\nelse:\n    device = torch.device('cpu')\n    \ndevice","metadata":{"papermill":{"duration":0.084185,"end_time":"2023-06-05T23:58:31.822492","exception":false,"start_time":"2023-06-05T23:58:31.738307","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:25.040024Z","iopub.execute_input":"2023-06-25T07:01:25.040364Z","iopub.status.idle":"2023-06-25T07:01:25.072228Z","shell.execute_reply.started":"2023-06-25T07:01:25.040319Z","shell.execute_reply":"2023-06-25T07:01:25.071013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config File","metadata":{"papermill":{"duration":0.006921,"end_time":"2023-06-05T23:58:31.836305","exception":false,"start_time":"2023-06-05T23:58:31.829384","status":"completed"},"tags":[]}},{"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 = False\n    if resize:\n        resize_value = 384\n        \n    # Model Settings    \n    model = 'UNET'\n    encoder = 'timm-resnest26d'\n    weights = 'imagenet'\n    \n    batch_size = 16\n    optimizer='Adam'\n    lr = 5e-4\n    epochs = 40","metadata":{"papermill":{"duration":0.01568,"end_time":"2023-06-05T23:58:31.858771","exception":false,"start_time":"2023-06-05T23:58:31.843091","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:25.074169Z","iopub.execute_input":"2023-06-25T07:01:25.074834Z","iopub.status.idle":"2023-06-25T07:01:25.081591Z","shell.execute_reply.started":"2023-06-25T07:01:25.0748Z","shell.execute_reply":"2023-06-25T07:01:25.080685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the Torch Dataset","metadata":{"papermill":{"duration":0.006619,"end_time":"2023-06-05T23:58:31.872455","exception":false,"start_time":"2023-06-05T23:58:31.865836","status":"completed"},"tags":[]}},{"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        self.normalize_image = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\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        # false color is scaled between 0 and 1!\n        return self.normalize_image(false_color).float(), human_pixel_mask.float()\n","metadata":{"papermill":{"duration":0.022272,"end_time":"2023-06-05T23:58:31.901509","exception":false,"start_time":"2023-06-05T23:58:31.879237","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:25.085034Z","iopub.execute_input":"2023-06-25T07:01:25.085671Z","iopub.status.idle":"2023-06-25T07:01:25.098085Z","shell.execute_reply.started":"2023-06-25T07:01:25.085633Z","shell.execute_reply":"2023-06-25T07:01:25.097246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the Training and Validation Dataloader","metadata":{"papermill":{"duration":0.006826,"end_time":"2023-06-05T23:58:31.915086","exception":false,"start_time":"2023-06-05T23:58:31.90826","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":2.233856,"end_time":"2023-06-05T23:58:34.156213","exception":false,"start_time":"2023-06-05T23:58:31.922357","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:25.100615Z","iopub.execute_input":"2023-06-25T07:01:25.101547Z","iopub.status.idle":"2023-06-25T07:01:25.474688Z","shell.execute_reply.started":"2023-06-25T07:01:25.101522Z","shell.execute_reply":"2023-06-25T07:01:25.473693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show some Images from the Dataloaders","metadata":{"papermill":{"duration":0.007255,"end_time":"2023-06-05T23:58:34.171163","exception":false,"start_time":"2023-06-05T23:58:34.163908","status":"completed"},"tags":[]}},{"cell_type":"code","source":"image,mask = next(iter(train_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":{"papermill":{"duration":5.573547,"end_time":"2023-06-05T23:58:39.752003","exception":false,"start_time":"2023-06-05T23:58:34.178456","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:25.476241Z","iopub.execute_input":"2023-06-25T07:01:25.476849Z","iopub.status.idle":"2023-06-25T07:01:30.760742Z","shell.execute_reply.started":"2023-06-25T07:01:25.476815Z","shell.execute_reply":"2023-06-25T07:01:30.759585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create the Model UNET","metadata":{"papermill":{"duration":0.011056,"end_time":"2023-06-05T23:58:39.774362","exception":false,"start_time":"2023-06-05T23:58:39.763306","status":"completed"},"tags":[]}},{"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))","metadata":{"execution":{"iopub.status.busy":"2023-06-25T07:01:30.762213Z","iopub.execute_input":"2023-06-25T07:01:30.762895Z","iopub.status.idle":"2023-06-25T07:01:39.85673Z","shell.execute_reply.started":"2023-06-25T07:01:30.762855Z","shell.execute_reply":"2023-06-25T07:01:39.854606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimizer","metadata":{"papermill":{"duration":0.012272,"end_time":"2023-06-05T23:58:45.624759","exception":false,"start_time":"2023-06-05T23:58:45.612487","status":"completed"},"tags":[]}},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=CFG.lr)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min',patience = 4, factor = 0.31622776601, verbose = True)\nprint(f'learning rate: {optimizer.param_groups[0][\"lr\"]}')","metadata":{"papermill":{"duration":0.025147,"end_time":"2023-06-05T23:58:45.663062","exception":false,"start_time":"2023-06-05T23:58:45.637915","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:39.858442Z","iopub.execute_input":"2023-06-25T07:01:39.85893Z","iopub.status.idle":"2023-06-25T07:01:39.869198Z","shell.execute_reply.started":"2023-06-25T07:01:39.858892Z","shell.execute_reply":"2023-06-25T07:01:39.867872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss Function","metadata":{"papermill":{"duration":0.012668,"end_time":"2023-06-05T23:58:45.688127","exception":false,"start_time":"2023-06-05T23:58:45.675459","status":"completed"},"tags":[]}},{"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","metadata":{"papermill":{"duration":0.021971,"end_time":"2023-06-05T23:58:45.756632","exception":false,"start_time":"2023-06-05T23:58:45.734661","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:39.870776Z","iopub.execute_input":"2023-06-25T07:01:39.871599Z","iopub.status.idle":"2023-06-25T07:01:39.879443Z","shell.execute_reply.started":"2023-06-25T07:01:39.871562Z","shell.execute_reply":"2023-06-25T07:01:39.878244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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)\n","metadata":{"papermill":{"duration":0.021768,"end_time":"2023-06-05T23:58:45.722462","exception":false,"start_time":"2023-06-05T23:58:45.700694","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:39.882979Z","iopub.execute_input":"2023-06-25T07:01:39.883721Z","iopub.status.idle":"2023-06-25T07:01:39.889782Z","shell.execute_reply.started":"2023-06-25T07:01:39.883693Z","shell.execute_reply":"2023-06-25T07:01:39.888615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training and Validation Loop","metadata":{"papermill":{"duration":0.012211,"end_time":"2023-06-05T23:58:45.781153","exception":false,"start_time":"2023-06-05T23:58:45.768942","status":"completed"},"tags":[]}},{"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 >=10:\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":{"papermill":{"duration":10491.737356,"end_time":"2023-06-06T02:53:37.531076","exception":false,"start_time":"2023-06-05T23:58:45.79372","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-25T07:01:39.891352Z","iopub.execute_input":"2023-06-25T07:01:39.892146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training and Validation History","metadata":{"papermill":{"duration":0.020762,"end_time":"2023-06-06T02:53:37.573583","exception":false,"start_time":"2023-06-06T02:53:37.552821","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.418136,"end_time":"2023-06-06T02:53:38.012961","exception":false,"start_time":"2023-06-06T02:53:37.594825","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show some predictions for the validation dataset","metadata":{"papermill":{"duration":0.021965,"end_time":"2023-06-06T02:53:38.056942","exception":false,"start_time":"2023-06-06T02:53:38.034977","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":15.675743,"end_time":"2023-06-06T02:53:53.754462","exception":false,"start_time":"2023-06-06T02:53:38.078719","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Bonus: Vary the threshold for the predictions","metadata":{"papermill":{"duration":0.088432,"end_time":"2023-06-06T02:53:53.932885","exception":false,"start_time":"2023-06-06T02:53:53.844453","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":260.664905,"end_time":"2023-06-06T02:58:14.685661","exception":false,"start_time":"2023-06-06T02:53:54.020756","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]}]}