{"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":"code","source":"!pip install monai","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:18.711166Z","iopub.execute_input":"2021-05-22T09:44:18.711482Z","iopub.status.idle":"2021-05-22T09:44:24.163480Z","shell.execute_reply.started":"2021-05-22T09:44:18.711455Z","shell.execute_reply":"2021-05-22T09:44:24.162567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport json\nimport csv\nimport random\nimport pickle\nimport cv2\nimport numpy as np\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\n\nimport torchvision.transforms as transforms\n\nfrom PIL import Image, ImageOps, ImageEnhance\nfrom torch.utils.data import Dataset, DataLoader\nfrom scipy.ndimage.measurements import label, center_of_mass\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.metrics import roc_auc_score, roc_curve\nplt.style.use('ggplot')\n","metadata":{"execution":{"iopub.status.busy":"2021-05-22T10:25:03.031195Z","iopub.execute_input":"2021-05-22T10:25:03.031615Z","iopub.status.idle":"2021-05-22T10:25:03.041493Z","shell.execute_reply.started":"2021-05-22T10:25:03.031577Z","shell.execute_reply":"2021-05-22T10:25:03.040215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.transforms import (\n    Activations,\n    AddChanneld,\n    AsDiscrete,\n    Compose,\n    CropForeground,\n    CropForegroundd,\n    LoadImaged,\n    NormalizeIntensityd,\n    RandAffined,\n    RandRotate90d,\n    RandCropByPosNegLabeld,\n    RandFlipd,\n    RandFlip,\n    RandRotated,\n    RandRotate,\n    Resized,\n    Resize,\n    ScaleIntensityd,\n    SpatialCropd,\n    SpatialCrop,\n    ToTensord,    \n)","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:24.178089Z","iopub.execute_input":"2021-05-22T09:44:24.178536Z","iopub.status.idle":"2021-05-22T09:44:24.188300Z","shell.execute_reply.started":"2021-05-22T09:44:24.178500Z","shell.execute_reply":"2021-05-22T09:44:24.187361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset class","metadata":{}},{"cell_type":"code","source":"class RefugeDataset(Dataset):\n\n    def __init__(self, \n                 root_dir, \n                 find_center_net,\n                 index_path = None,\n                 roi_size=600,\n                 split='train', \n                 data_augm=None\n                ):\n        # Define attributes\n        self.root_dir = root_dir\n        self.split = split\n        if split != 'test':\n            self.transform = Compose(\n                [               \n                    CropForegroundd(keys=[\"img\",\"od\", \"oc\"], source_key=\"img\"),\n                    Resized(keys=[\"img\",\"od\", \"oc\"], \n                            spatial_size=[IMG_H, IMG_W], \n                            mode=('bilinear', 'nearest', 'nearest')),\n                ]\n            )\n        else:\n            self.transform = Compose(\n            [               \n                CropForegroundd(keys=[\"img\"], source_key=\"img\"),\n                Resized(keys=[\"img\"], \n                        spatial_size=[IMG_H, IMG_W], \n                        mode=('bilinear')),\n            ]\n        )\n        self.resizer_img = Resize(spatial_size=[RESIZE_W, RESIZE_H], mode=\"bilinear\")\n        self.resizer_seg = Resize(spatial_size=[RESIZE_W, RESIZE_H], mode='nearest')\n\n        self.data_augm = data_augm\n        # Load data index\n        if not index_path:\n            index_path = os.path.join(self.root_dir, self.split, 'index.json')\n        with open(index_path) as f:\n            #index = json.load(f)\n#             self.index = {'0': index['0'], \n#                           '1': index['1'], \n#                           '2': index['300'], \n#                           '3': index['200'], \n#                           '4': index['100']\n#                         }\n            self.index = json.load(f)\n            \n        self.images = []\n        self.segs = []\n        if split != 'test':\n            for k in range(len(self.index)):\n                print('Loading {} image {}/{}...'.format(split, k, len(self.index)), end='\\r')\n                img_name = os.path.join(self.root_dir, self.split, 'images', self.index[str(k)]['ImgName'])\n                img = np.array(Image.open(img_name).convert('RGB'))\n                img = transforms.functional.to_tensor(img).numpy()\n                seg_name = os.path.join(self.root_dir, self.split, 'gts', self.index[str(k)]['ImgName'].split('.')[0]+'.bmp')\n                seg = np.array(Image.open(seg_name)).copy()\n                seg = 255. - seg\n                od = (seg>=127.).astype(np.float32)\n                oc = (seg>=250.).astype(np.float32)\n                od = transforms.functional.to_tensor(od).numpy()\n                oc = transforms.functional.to_tensor(oc).numpy()\n\n                data = {'img': img, 'od': od, 'oc': oc}\n                data = self.transform(data)\n                seg = np.concatenate([data['od'], data['oc']], axis=0)\n                img = data['img']\n                \n                img, seg = center_crop_and_resize(find_center_net, img, seg, \n                                                  roi_size, self.resizer_img, \n                                                  self.resizer_seg)\n#                 img = np.moveaxis(img, 0, 2)\n#                 img = Image.fromarray(np.uint8(img*255))\n#                 img = ImageOps.autocontrast(img, cutoff=0, ignore=0)\n#                 img = np.array(img)\n#                 img = transforms.functional.to_tensor(img).numpy()\n                \n                self.images.append(img)\n                self.segs.append(seg)\n        else:\n            for k in range(len(self.index)):\n                print('Loading {} image {}/{}...'.format(split, k, len(self.index)), end='\\r')\n                img_name = os.path.join(self.root_dir, self.split, 'images', self.index[str(k)]['ImgName'])\n                img = np.array(Image.open(img_name).convert('RGB'))\n                img = transforms.functional.to_tensor(img).numpy()\n                data = {'img': img}\n                data = self.transform(data)\n                \n                img, seg = center_crop_and_resize(find_center_net, data['img'], None, \n                                                  roi_size, self.resizer_img, \n                                                  None)\n#                 img = np.moveaxis(img, 0, 2)\n#                 img = Image.fromarray(np.uint8(img*255))\n#                 img = ImageOps.autocontrast(img, cutoff=0, ignore=0)\n#                 img = np.array(img)\n#                 img = transforms.functional.to_tensor(img).numpy()\n                \n                self.images.append(img)\n        print('Succesfully loaded {} dataset.'.format(split) + ' '*50)\n            \n            \n    def __len__(self):\n        return len(self.index)\n\n    def __getitem__(self, idx):\n        # Image\n        img = self.images[idx]\n    \n        # Return only images for 'test' set\n        if self.split == 'test':\n            return torch.Tensor(img)\n        \n        # Else, images and ground truth\n        else:\n            # Label\n            lab = torch.tensor(self.index[str(idx)]['Label'], dtype=torch.float32)\n\n            # Segmentation masks\n            seg = self.segs[idx]\n\n            # Fovea localization\n            f_x = self.index[str(idx)]['Fovea_X']\n            f_y = self.index[str(idx)]['Fovea_Y']\n            fov = torch.FloatTensor([f_x, f_y])\n        \n            if self.split == 'train' and self.data_augm != None:                \n#                 img = np.moveaxis(img, 0, 2)\n#                 img = Image.fromarray(np.uint8(img*255))\n                \n#                 brightness_range = [0.9, 1.1]\n#                 brightness_k = random.uniform(brightness_range[0], brightness_range[1])\n#                 img = ImageEnhance.Brightness(img)\n#                 img = img.enhance(brightness_k)\n#                 img = np.array(img)\n#                 img = transforms.functional.to_tensor(img).numpy()\n\n                data = self.data_augm({'img': img, 'seg': seg})\n                img, seg = data['img'], data['seg']\n            return torch.Tensor(img), lab, torch.Tensor(seg),  \\\n                    fov, self.index[str(idx)]['ImgName']","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:24.189835Z","iopub.execute_input":"2021-05-22T09:44:24.190265Z","iopub.status.idle":"2021-05-22T09:44:24.216338Z","shell.execute_reply.started":"2021-05-22T09:44:24.190228Z","shell.execute_reply":"2021-05-22T09:44:24.215183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics","metadata":{}},{"cell_type":"code","source":"EPS = 1e-7\n\ndef compute_dice_coef(input, target):\n    '''\n    Compute dice score metric.\n    '''\n    batch_size = input.shape[0]\n    return sum([dice_coef_sample(input[k,:,:], target[k,:,:]) for k in range(batch_size)])/batch_size\n\ndef dice_coef_sample(input, target):\n    iflat = input.contiguous().view(-1)\n    tflat = target.contiguous().view(-1)\n    intersection = (iflat * tflat).sum()\n    return (2. * intersection) / (iflat.sum() + tflat.sum())\n\n\ndef vertical_diameter(binary_segmentation):\n    '''\n    Get the vertical diameter from a binary segmentation.\n    The vertical diameter is defined as the \"fattest\" area of the binary_segmentation parameter.\n    '''\n\n    # get the sum of the pixels in the vertical axis\n    vertical_axis_diameter = np.sum(binary_segmentation, axis=1)\n\n    # pick the maximum value\n    diameter = np.max(vertical_axis_diameter, axis=1)\n\n    # return it\n    return diameter\n\n\n\ndef vertical_cup_to_disc_ratio(od, oc):\n    '''\n    Compute the vertical cup-to-disc ratio from a given labelling map.\n    '''\n    # compute the cup diameter\n    cup_diameter = vertical_diameter(oc)\n    # compute the disc diameter\n    disc_diameter = vertical_diameter(od)\n\n    return cup_diameter / (disc_diameter + EPS)\n\ndef compute_vCDR_error(pred_od, pred_oc, gt_od, gt_oc):\n    '''\n    Compute vCDR prediction error, along with predicted vCDR and ground truth vCDR.\n    '''\n    pred_vCDR = vertical_cup_to_disc_ratio(pred_od, pred_oc)\n    gt_vCDR = vertical_cup_to_disc_ratio(gt_od, gt_oc)\n    vCDR_err = np.mean(np.abs(gt_vCDR - pred_vCDR))\n    return vCDR_err, pred_vCDR, gt_vCDR\n\n\ndef classif_eval(classif_preds, classif_gts):\n    '''\n    Compute AUC classification score.\n    '''\n    auc = roc_auc_score(classif_gts, classif_preds)\n    return auc\n\n\ndef fov_error(pred_fov, gt_fov):\n    '''\n    Fovea localization error metric (mean root squared error).\n    '''\n    err = np.sqrt(np.sum((gt_fov-pred_fov)**2, axis=1)).mean()\n    return err","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:26.406337Z","iopub.execute_input":"2021-05-22T09:44:26.406655Z","iopub.status.idle":"2021-05-22T09:44:26.416894Z","shell.execute_reply.started":"2021-05-22T09:44:26.406623Z","shell.execute_reply":"2021-05-22T09:44:26.416014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Post-processing functions","metadata":{}},{"cell_type":"code","source":"def refine_seg(pred):\n    '''\n    Only retain the biggest connected component of a segmentation map.\n    '''\n    np_pred = pred.numpy()\n        \n    largest_ccs = []\n    for i in range(np_pred.shape[0]):\n        labeled, ncomponents = label(np_pred[i,:,:])\n        bincounts = np.bincount(labeled.flat)[1:]\n        if len(bincounts) == 0:\n            largest_cc = labeled == 0\n        else:\n            largest_cc = labeled == np.argmax(bincounts)+1\n        largest_cc = torch.tensor(largest_cc, dtype=torch.float32)\n        largest_ccs.append(largest_cc)\n    largest_ccs = torch.stack(largest_ccs)\n    \n    return largest_ccs","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:31.088713Z","iopub.execute_input":"2021-05-22T09:44:31.089060Z","iopub.status.idle":"2021-05-22T09:44:31.098179Z","shell.execute_reply.started":"2021-05-22T09:44:31.089029Z","shell.execute_reply":"2021-05-22T09:44:31.097132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Network","metadata":{}},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, n_channels=3, n_classes=2):\n        super(UNet, self).__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.epoch = 0\n\n        self.inc = DoubleConv(n_channels, 64)\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        factor = 2 \n        self.down4 = Down(512, 1024 // factor)\n        self.up1 = Up(1024, 512 // factor)\n        self.up2 = Up(512, 256 // factor)\n        self.up3 = Up(256, 128 // factor)\n        self.up4 = Up(128, 64)\n        self.output_layer = OutConv(64, n_classes)\n        #self.output_fc = FullyConnect(1024, 1)\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        out = self.up1(x5, x4)\n        out = self.up2(out, x3)\n        out = self.up3(out, x2)\n        out = self.up4(out, x1)\n        out = self.output_layer(out)\n        out = torch.sigmoid(out)\n    \n        return out\n\n    \nclass FullyConnect(nn.Module):\n    \n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool = nn.AdaptiveMaxPool2d(1)\n        self.fc = nn.Linear(in_channels, out_channels, bias=True)\n    \n    def forward(self, x):\n        x = self.maxpool(x).reshape(x.shape[0], -1)\n        out = self.fc(x)\n        return out\n        \nclass DoubleConv(nn.Module):\n    \"\"\"(convolution => [BN] => ReLU) * 2\"\"\"\n\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n\nclass Down(nn.Module):\n    \"\"\"Downscaling with maxpool then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n\n        # Use the normal convolutions to reduce the number of channels\n        self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)\n\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    '''\n    Simple convolution.\n    '''\n    def __init__(self, in_channels, out_channels):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:33.177845Z","iopub.execute_input":"2021-05-22T09:44:33.178206Z","iopub.status.idle":"2021-05-22T09:44:33.199348Z","shell.execute_reply.started":"2021-05-22T09:44:33.178173Z","shell.execute_reply":"2021-05-22T09:44:33.198010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Settings","metadata":{}},{"cell_type":"code","source":"root_dir = '/kaggle/input/eurecom-aml-2021-challenge-2/refuge_data/refuge_data'\nlr = 1e-4\nbatch_size = 8\nnum_workers = 8\ntotal_epoch = 100","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:35.426471Z","iopub.execute_input":"2021-05-22T09:44:35.426789Z","iopub.status.idle":"2021-05-22T09:44:35.433178Z","shell.execute_reply.started":"2021-05-22T09:44:35.426760Z","shell.execute_reply":"2021-05-22T09:44:35.432163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing transforms","metadata":{}},{"cell_type":"code","source":"# img and seg_gt as numpy\ndef center_crop_and_resize(model, img, seg_gt, roi_size, img_resizer, seg_resizer):   \n    img_res = img_resizer(img)\n    logits = model(torch.Tensor(img_res).unsqueeze(0).to(device))\n    center = find_mass_center((logits[:,0,:,:]>=0.5).type(torch.int8).cpu(), \n                              img.shape[1:3])\n    \n    cropper = SpatialCrop(roi_center=center, roi_size=roi_size)\n    img = cropper(img)\n    img = img_resizer(img)\n    if seg_resizer != None:\n        seg_gt = cropper(seg_gt)\n        seg_gt = seg_resizer(seg_gt)\n    return img, seg_gt\n\ndef find_mass_center(seg, original_shape):\n    current_shape = seg.shape[1:3]\n    largest_ccs = refine_seg(seg).numpy()\n\n    center = center_of_mass(largest_ccs[0])\n\n    original_center = (center[0] * original_shape[0] / current_shape[0],\n                       center[1] * original_shape[1] / current_shape[1])\n    return original_center\n\n\ndef crop_data(img, center, roi_size):\n    cropped_imgs = torch.zeros([img.shape[0], roi_size, roi_size])\n    cropped_segs = torch.zeros([img.shape[0], roi_size, roi_size])\n    return cropped_imgs","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:38.496887Z","iopub.execute_input":"2021-05-22T09:44:38.497236Z","iopub.status.idle":"2021-05-22T09:44:38.506610Z","shell.execute_reply.started":"2021-05-22T09:44:38.497206Z","shell.execute_reply":"2021-05-22T09:44:38.505469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROI_SIZE = 600\nIMG_W = 1634\nIMG_H = 1634\nRESIZE_W = 256\nRESIZE_H = 256\n\n\n# Data augmentation function, called in dataset.__get_item__\n# Requires keys=['img', 'seg']\ntrain_data_augm = Compose(\n            [        \n                RandFlipd(keys=['img', 'seg'], prob=0.5, spatial_axis=1),\n                #RandFlipd(keys=['img', 'seg'], prob=0.5, spatial_axis=0),\n                #RandRotated(keys=['img', 'seg'], range_x=10, prob=0.5),\n                ToTensord(keys=['img', 'seg'])\n            ]\n        )","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:44.489512Z","iopub.execute_input":"2021-05-22T09:44:44.489846Z","iopub.status.idle":"2021-05-22T09:44:44.498772Z","shell.execute_reply.started":"2021-05-22T09:44:44.489814Z","shell.execute_reply":"2021-05-22T09:44:44.498023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Device, model, loss and optimizer","metadata":{}},{"cell_type":"code","source":"# Device\ndevice = torch.device(\"cuda:0\")\n\n# Find center Net\nfind_center_net = UNet(n_channels=3, n_classes=2).to(device)\nfind_center_net.load_state_dict(torch.load('/kaggle/input/unet-pretrained/UNet_9557_9103_fgcrop.pth'))\n\n\n# Network\nmodel = UNet(n_channels=3, n_classes=2).to(device)\n\n# Loss\nseg_loss = torch.nn.BCELoss(reduction='mean')\n\n# Optimizer\noptimizer = optim.Adam(model.parameters(), lr=lr)\nscheduler = lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5)","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:51.463396Z","iopub.execute_input":"2021-05-22T09:44:51.463747Z","iopub.status.idle":"2021-05-22T09:44:51.871975Z","shell.execute_reply.started":"2021-05-22T09:44:51.463717Z","shell.execute_reply":"2021-05-22T09:44:51.871184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create datasets and data loaders\nAll image files are loaded in RAM in order to speed up the pipeline. Therefore, each dataset creation should take a few minutes.","metadata":{}},{"cell_type":"code","source":"# Datasets\ntrain_set = RefugeDataset(root_dir, \n                          split='train',\n                          index_path = \"/kaggle/input/unet-pretrained/balanced_index.json\",\n                          find_center_net=find_center_net,\n                          data_augm=train_data_augm,\n                          roi_size=ROI_SIZE)\nval_set = RefugeDataset(root_dir, \n                        split='val',\n                        find_center_net=find_center_net,\n                        roi_size=ROI_SIZE)\ntest_set = RefugeDataset(root_dir, \n                        split='test',\n                        find_center_net=find_center_net,\n                        roi_size=ROI_SIZE)\n\n# Dataloaders\ntrain_loader = DataLoader(train_set, \n                          batch_size=batch_size, \n                          shuffle=True, \n                          num_workers=num_workers,\n                          pin_memory=True,\n                         )\nval_loader = DataLoader(val_set, \n                        batch_size=batch_size, \n                        shuffle=False, \n                        num_workers=num_workers,\n                        pin_memory=True,\n                        )\ntest_loader = DataLoader(test_set, \n                       batch_size=batch_size, \n                       shuffle=False, \n                       num_workers=num_workers,\n                       pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:44:58.227202Z","iopub.execute_input":"2021-05-22T09:44:58.227539Z","iopub.status.idle":"2021-05-22T09:52:20.837241Z","shell.execute_reply.started":"2021-05-22T09:44:58.227508Z","shell.execute_reply":"2021-05-22T09:52:20.835473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ","metadata":{}},{"cell_type":"markdown","source":"# Train for OC/OD segmentation","metadata":{}},{"cell_type":"code","source":"# Define parameters\nnb_train_batches = len(train_loader)\nnb_val_batches = len(val_loader)\nnb_iter = 0\nbest_val_auc = 0.\n\nwhile model.epoch < total_epoch:\n    # Accumulators\n    train_vCDRs, val_vCDRs = [], []\n    train_classif_gts, val_classif_gts = [], []\n    train_loss, val_loss = 0., 0.\n    train_dsc_od, val_dsc_od = 0., 0.\n    train_dsc_oc, val_dsc_oc = 0., 0.\n    train_vCDR_error, val_vCDR_error = 0., 0.\n    \n    ############\n    # TRAINING #\n    ############\n    model.train()\n    train_data = iter(train_loader)\n    for k in range(nb_train_batches):\n        # Loads data\n        imgs, classif_gts, seg_gts, fov_coords, names = train_data.next()\n        imgs, classif_gts, seg_gts = imgs.to(device), classif_gts.to(device), seg_gts.to(device)\n\n        # Forward pass\n        logits = model(imgs)\n        loss = seg_loss(logits, seg_gts)\n \n        # Backward pass\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        train_loss += loss.item() / nb_train_batches\n        \n        with torch.no_grad():\n            # Compute segmentation metric\n            pred_od = refine_seg((logits[:,0,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n            pred_oc = refine_seg((logits[:,1,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n            gt_od = seg_gts[:,0,:,:].type(torch.int8)\n            gt_oc = seg_gts[:,1,:,:].type(torch.int8)\n            dsc_od = compute_dice_coef(pred_od, gt_od)\n            dsc_oc = compute_dice_coef(pred_oc, gt_oc)\n            train_dsc_od += dsc_od.item()/nb_train_batches\n            train_dsc_oc += dsc_oc.item()/nb_train_batches\n\n\n            # Compute and store vCDRs\n            vCDR_error, pred_vCDR, gt_vCDR = compute_vCDR_error(pred_od.cpu().numpy(), pred_oc.cpu().numpy(), gt_od.cpu().numpy(), gt_oc.cpu().numpy())\n            train_vCDRs += pred_vCDR.tolist()\n            train_vCDR_error += vCDR_error / nb_train_batches\n            train_classif_gts += classif_gts.cpu().numpy().tolist()\n            \n        # Increase iterations\n        nb_iter += 1\n        \n        # Std out\n        print('Epoch {}, iter {}/{}, loss {:.6f}'.format(model.epoch+1, k+1, nb_train_batches, loss.item()) + ' '*20, \n              end='\\r')\n        \n    # Train a logistic regression on vCDRs\n    train_vCDRs = np.array(train_vCDRs).reshape(-1,1)\n    train_classif_gts = np.array(train_classif_gts)\n    clf = LogisticRegression(random_state=0, solver='lbfgs').fit(train_vCDRs, train_classif_gts)\n    train_classif_preds = clf.predict_proba(train_vCDRs)[:,1]\n    train_auc = classif_eval(train_classif_preds, train_classif_gts)\n    \n    ##############\n    # VALIDATION #\n    ##############\n    model.eval()\n    with torch.no_grad():\n        val_data = iter(val_loader)\n        for k in range(nb_val_batches):\n            # Loads data\n            imgs, classif_gts, seg_gts, fov_coords, names = val_data.next()\n            imgs, classif_gts, seg_gts = imgs.to(device), classif_gts.to(device), seg_gts.to(device)\n\n            # Forward pass\n            logits = model(imgs)\n            val_loss += seg_loss(logits, seg_gts).item() / nb_val_batches\n\n            # Std out\n            print('Validation iter {}/{}'.format(k+1, nb_val_batches) + ' '*50, \n                  end='\\r')\n            \n            # Compute segmentation metric\n            pred_od = refine_seg((logits[:,0,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n            pred_oc = refine_seg((logits[:,1,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n            gt_od = seg_gts[:,0,:,:].type(torch.int8)\n            gt_oc = seg_gts[:,1,:,:].type(torch.int8)\n            dsc_od = compute_dice_coef(pred_od, gt_od)\n            dsc_oc = compute_dice_coef(pred_oc, gt_oc)\n            val_dsc_od += dsc_od.item()/nb_val_batches\n            val_dsc_oc += dsc_oc.item()/nb_val_batches\n            \n            # Compute and store vCDRs\n            vCDR_error, pred_vCDR, gt_vCDR = compute_vCDR_error(pred_od.cpu().numpy(), pred_oc.cpu().numpy(), gt_od.cpu().numpy(), gt_oc.cpu().numpy())\n            val_vCDRs += pred_vCDR.tolist()\n            val_vCDR_error += vCDR_error / nb_val_batches\n            val_classif_gts += classif_gts.cpu().numpy().tolist()\n            \n\n    # Glaucoma predictions from vCDRs\n    val_vCDRs = np.array(val_vCDRs).reshape(-1,1)\n    val_classif_gts = np.array(val_classif_gts)\n    val_classif_preds = clf.predict_proba(val_vCDRs)[:,1]\n    val_auc = classif_eval(val_classif_preds, val_classif_gts)\n        \n    # Validation results\n    print('VALIDATION epoch {}'.format(model.epoch+1)+' '*50)\n    print('LOSSES: {:.4f} (train), {:.4f} (val)'.format(train_loss, val_loss))\n    print('OD segmentation (Dice Score): {:.4f} (train), {:.4f} (val)'.format(train_dsc_od, val_dsc_od))\n    print('OC segmentation (Dice Score): {:.4f} (train), {:.4f} (val)'.format(train_dsc_oc, val_dsc_oc))\n    print('vCDR error: {:.4f} (train), {:.4f} (val)'.format(train_vCDR_error, val_vCDR_error))\n    print('Classification (AUC): {:.4f} (train), {:.4f} (val)'.format(train_auc, val_auc))\n    \n    # Save model if best validation AUC is reached\n    if val_auc > best_val_auc:\n        torch.save(model.state_dict(), '/kaggle/working/best_AUC_weights.pth')\n        with open('/kaggle/working/best_AUC_classifier.pkl', 'wb') as clf_file:\n            pickle.dump(clf, clf_file)\n        best_val_auc = val_auc\n        print('Best validation AUC reached. Saved model weights and classifier.')\n    print('_'*50)\n        \n    # End of epoch\n    model.epoch += 1\n    scheduler.step( (val_dsc_od * val_dsc_oc) * 0.5 )\n    print('CURRENT LR {:.8f}'.format(optimizer.param_groups[0]['lr']))\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load best model + classifier","metadata":{}},{"cell_type":"code","source":"# Load model and classifier\nmodel_val = UNet(n_channels=3, n_classes=2).to(device)\nmodel_val.load_state_dict(torch.load('/kaggle/working/best_AUC_weights.pth'))\nwith open('/kaggle/working/best_AUC_classifier.pkl', 'rb') as clf_file:\n   clf_val = pickle.load(clf_file)\n# model_val.load_state_dict(torch.load('../input/unet-pretrained/models/best_AUC_weights_2.pth'))\n# with open('../input/unet-pretrained/models/best_AUC_classifier_2.pkl', 'rb') as clf_file:\n#     clf_val = pickle.load(clf_file)","metadata":{"execution":{"iopub.status.busy":"2021-05-22T10:21:07.116877Z","iopub.execute_input":"2021-05-22T10:21:07.117299Z","iopub.status.idle":"2021-05-22T10:21:07.385243Z","shell.execute_reply.started":"2021-05-22T10:21:07.117253Z","shell.execute_reply":"2021-05-22T10:21:07.384401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Check performance is maintained on validation","metadata":{}},{"cell_type":"code","source":"model_val.eval()\nval_vCDRs = []\nval_classif_gts = []\nval_loss = 0.\nval_dsc_od = 0.\nval_dsc_oc = 0.\nval_vCDR_error = 0.\nnb_val_batches = len(val_loader)\nnb_iter = 0\n\nwith torch.no_grad():\n    val_data = iter(val_loader)\n    for k in range(nb_val_batches):\n        # Loads data\n        imgs, classif_gts, seg_gts, fov_coords, names = val_data.next()\n        imgs, classif_gts, seg_gts = imgs.to(device), classif_gts.to(device), seg_gts.to(device)\n\n        # Forward pass\n        logits = model_val(imgs)\n        val_loss += seg_loss(logits, seg_gts).item() / nb_val_batches\n\n        # Std out\n        print('Validation iter {}/{}'.format(k+1, nb_val_batches) + ' '*50, \n              end='\\r')\n\n        # Compute segmentation metric\n        pred_od = refine_seg((logits[:,0,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n        pred_oc = refine_seg((logits[:,1,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n        gt_od = seg_gts[:,0,:,:].type(torch.int8)\n        gt_oc = seg_gts[:,1,:,:].type(torch.int8)\n        dsc_od = compute_dice_coef(pred_od, gt_od)\n        dsc_oc = compute_dice_coef(pred_oc, gt_oc)\n        val_dsc_od += dsc_od.item()/nb_val_batches\n        val_dsc_oc += dsc_oc.item()/nb_val_batches\n        \n        # Compute and store vCDRs\n        vCDR_error, pred_vCDR, gt_vCDR = compute_vCDR_error(pred_od.cpu().numpy(), pred_oc.cpu().numpy(), gt_od.cpu().numpy(), gt_oc.cpu().numpy())\n        val_vCDRs += pred_vCDR.tolist()\n        val_vCDR_error += vCDR_error / nb_val_batches\n        val_classif_gts += classif_gts.cpu().numpy().tolist()\n\n\n# Glaucoma predictions from vCDRs\nval_vCDRs = np.array(val_vCDRs).reshape(-1,1)\nval_classif_gts = np.array(val_classif_gts)\nval_classif_preds = clf_val.predict_proba(val_vCDRs)[:,1]\nval_auc = classif_eval(val_classif_preds, val_classif_gts)\n\n# Validation results\nprint('VALIDATION '+' '*50)\nprint('LOSSES: {:.4f} (val)'.format(val_loss))\nprint('OD segmentation (Dice Score): {:.4f} (val)'.format(val_dsc_od))\nprint('OC segmentation (Dice Score): {:.4f} (val)'.format(val_dsc_oc))\nprint('vCDR error: {:.4f} (val)'.format(val_vCDR_error))\nprint('Classification (AUC): {:.4f} (val)'.format(val_auc))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_label(true_label, pred_label, img):\n    mtl = np.ma.masked_where(true_label <= 0.01, true_label)\n    mpl = np.ma.masked_where(pred_label <= 0.1, pred_label)\n\n    fig, ax = plt.subplots(1,2, figsize=(10,10))\n    ax[0].imshow(img, cmap=plt.get_cmap('binary'))\n    ax[0].imshow(mtl[0], cmap=plt.get_cmap('jet'), alpha=0.5, vmin=0, vmax=1)\n    ax[0].imshow(mtl[1], cmap=plt.get_cmap('Blues'), alpha=0.8, vmin=0, vmax=1)\n\n    ax[0].set_title(\"Ground truth\")\n    ax[0].grid(False)\n    ax[1].imshow(img, cmap=plt.get_cmap('binary'))\n    ax[1].imshow(mpl[0], cmap=plt.get_cmap('jet'), alpha=0.5, vmin=0, vmax=1)\n    ax[1].imshow(mpl[1], cmap=plt.get_cmap('Blues'), alpha=0.8, vmin=0, vmax=1)\n    ax[1].set_title(\"Prediction\")\n    ax[1].grid(False)\n    plt.savefig('mask_comparison.jpg')\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2021-05-22T10:26:30.433742Z","iopub.execute_input":"2021-05-22T10:26:30.434084Z","iopub.status.idle":"2021-05-22T10:26:30.443205Z","shell.execute_reply.started":"2021-05-22T10:26:30.434055Z","shell.execute_reply":"2021-05-22T10:26:30.442325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nidx = 333\nx = val_set.images[idx]\nx = np.moveaxis(x, 0, 2)\n#plt.imshow(x)\nlogits = model_val(torch.from_numpy(val_set.images[idx]).unsqueeze(0).to(device))\npred_od = refine_seg((logits[:,0,:,:]>=0.5).type(torch.int8).cpu()).to(device)\npred_oc = refine_seg((logits[:,1,:,:]>=0.5).type(torch.int8).cpu()).to(device)\npred = torch.vstack((pred_od, pred_oc)).cpu().detach().numpy()\nprint(pred.shape)\nlab = val_set.segs[idx]\nimg = np.moveaxis(val_set.images[idx], 0, 2)\nplot_label(lab, pred, img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot  as plt\ndef plot_ROC(classif_preds, classif_gts):\n    fpr, tpr, _ = roc_curve(classif_gts, classif_preds)\n    plt.title('Receiver Operating Characteristic')\n    plt.fill_between(fpr,np.zeros_like(tpr),tpr,alpha=0.2)\n    plt.plot(fpr, tpr)\n    plt.plot([0, 1], ls=\"--\")\n    plt.plot([0, 0], [1, 0] , c=\".7\"), plt.plot([1, 1] , c=\".7\")\n    plt.ylabel('True Positive Rate')\n    plt.xlabel('False Positive Rate')\n    plt.savefig('roc.jpg')\n    plt.show()\n    \nplot_ROC(val_classif_preds, val_classif_gts)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predictions on test set","metadata":{}},{"cell_type":"code","source":"nb_test_batches = len(test_loader)\nmodel_val.eval()\ntest_vCDRs = []\nwith torch.no_grad():\n    test_data = iter(test_loader)\n    for k in range(nb_test_batches):\n        # Loads data\n        imgs = test_data.next()\n        imgs = imgs.to(device)\n\n        # Forward pass\n        logits = model_val(imgs)\n\n        # Std out\n        print('Test iter {}/{}'.format(k+1, nb_test_batches) + ' '*50, \n              end='\\r')\n            \n        # Compute segmentation\n        pred_od = refine_seg((logits[:,0,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n        pred_oc = refine_seg((logits[:,1,:,:]>=0.5).type(torch.int8).cpu()).to(device)\n        # Compute and store vCDRs\n        pred_vCDR = vertical_cup_to_disc_ratio(pred_od.cpu().numpy(), pred_oc.cpu().numpy())\n        test_vCDRs += pred_vCDR.tolist()\n            \n\n    # Glaucoma predictions from vCDRs\n    test_vCDRs = np.array(test_vCDRs).reshape(-1,1)\n    test_classif_preds = clf_val.predict_proba(test_vCDRs)[:,1]\n    \n# Prepare and save .csv file\ndef create_submission_csv(prediction, submission_filename='/kaggle/working/submission.csv'):\n    \"\"\"Create a sumbission file in the appropriate format for evaluation.\n\n    :param\n    prediction: list of predictions (ex: [0.12720, 0.89289, ..., 0.29829])\n    \"\"\"\n    \n    with open(submission_filename, mode='w') as csv_file:\n        fieldnames = ['Id', 'Predicted']\n        writer = csv.DictWriter(csv_file, fieldnames=fieldnames)\n        writer.writeheader()\n\n        for i, p in enumerate(prediction):\n            writer.writerow({'Id': \"T{:04d}\".format(i+1), 'Predicted': '{:f}'.format(p)})\n\ncreate_submission_csv(test_classif_preds)\ncreate_submission_csv(val_classif_preds, submission_filename=\"/kaggle/working/validation.csv\")\n# The submission.csv file is under /kaggle/working/submission.csv.\n# If you want to submit it, you should download it before closing the current kernel.","metadata":{"execution":{"iopub.status.busy":"2021-05-22T09:26:56.554472Z","iopub.execute_input":"2021-05-22T09:26:56.554800Z","iopub.status.idle":"2021-05-22T09:27:01.695985Z","shell.execute_reply.started":"2021-05-22T09:26:56.554768Z","shell.execute_reply":"2021-05-22T09:27:01.694380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}