{"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":"import os\nimport shutil\nimport json\nimport csv\nimport random\nimport pickle\nimport cv2\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torchvision.transforms as transforms\n\n\nimport PIL\nfrom PIL import Image, ImageOps\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\nfrom scipy.ndimage.measurements import label\nfrom sklearn.linear_model import LogisticRegression\nfrom sklearn.metrics import roc_auc_score, roc_curve\n","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:25.342446Z","iopub.execute_input":"2021-05-23T19:53:25.342763Z","iopub.status.idle":"2021-05-23T19:53:25.349228Z","shell.execute_reply.started":"2021-05-23T19:53:25.342732Z","shell.execute_reply":"2021-05-23T19:53:25.348156Z"},"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, root_dir, split='train', output_size=(256,256)):\n        # Define attributes\n        self.output_size = output_size\n        self.root_dir = root_dir\n        self.split = split\n        \n        # Load data index\n        with open(os.path.join(self.root_dir, self.split, 'index.json')) as f:\n            self.index = json.load(f)\n            \n        self.images = []\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)]['IMG_NAME'])\n            img = np.array(Image.open(img_name).convert('RGB'))\n            img = transforms.functional.to_tensor(img)\n            img = transforms.functional.resize(img, self.output_size, interpolation=Image.BILINEAR)\n            self.images.append(img)\n            \n        # Load ground truth for 'train' and 'val' sets\n        if split != 'test':\n            self.segs = []\n            for k in range(len(self.index)):\n                print('Loading {} segmentation {}/{}...'.format(split, k, len(self.index)), end='\\r')\n                seg_name = os.path.join(self.root_dir, self.split, 'gts', self.index[str(k)]['IMG_NAME'].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 = torch.from_numpy(od[None,:,:])\n                oc = torch.from_numpy(oc[None,:,:])\n                od = transforms.functional.resize(od, self.output_size, interpolation=Image.NEAREST)\n                oc = transforms.functional.resize(oc, self.output_size, interpolation=Image.NEAREST)\n                seg = torch.cat([od, oc], dim=0)\n                self.segs.append(seg)\n                \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 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            return img, lab, seg, fov, self.index[str(idx)]['IMG_NAME']","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:26.424952Z","iopub.execute_input":"2021-05-23T19:53:26.425266Z","iopub.status.idle":"2021-05-23T19:53:26.441941Z","shell.execute_reply.started":"2021-05-23T19:53:26.425237Z","shell.execute_reply":"2021-05-23T19:53:26.441012Z"},"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-23T19:53:27.661209Z","iopub.execute_input":"2021-05-23T19:53:27.661542Z","iopub.status.idle":"2021-05-23T19:53:27.672063Z","shell.execute_reply.started":"2021-05-23T19:53:27.661511Z","shell.execute_reply":"2021-05-23T19:53:27.671095Z"},"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-23T19:53:28.687347Z","iopub.execute_input":"2021-05-23T19:53:28.687690Z","iopub.status.idle":"2021-05-23T19:53:28.693952Z","shell.execute_reply.started":"2021-05-23T19:53:28.687660Z","shell.execute_reply":"2021-05-23T19:53:28.693062Z"},"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\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        return out\n    \nclass RN(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\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        return out\n\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-23T19:53:29.919786Z","iopub.execute_input":"2021-05-23T19:53:29.920095Z","iopub.status.idle":"2021-05-23T19:53:29.938991Z","shell.execute_reply.started":"2021-05-23T19:53:29.920067Z","shell.execute_reply":"2021-05-23T19:53:29.937928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Settings","metadata":{}},{"cell_type":"markdown","source":"## DATA AUGMENTATION","metadata":{}},{"cell_type":"code","source":"dir_path=\"../input/eurecom-aml-2021-challenge-2/refuge_data/refuge_data/\"\ndf_train = pd.read_json(dir_path+\"train/index.json\").T.rename(columns={\"ImgName\" : \"IMG_NAME\"})\ndf_val = pd.read_json(dir_path+\"val/index.json\").T.rename(columns={\"ImgName\" : \"IMG_NAME\"})\ndf_test = pd.read_json(dir_path+\"test/index.json\").T.rename(columns={\"ImgName\" : \"IMG_NAME\"})\n\nIMG_SIZE = 512\nNUM_CLASSES = 5\nSEED = 77\nTRAIN_NUM = 1000","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:31.819301Z","iopub.execute_input":"2021-05-23T19:53:31.819630Z","iopub.status.idle":"2021-05-23T19:53:32.515032Z","shell.execute_reply.started":"2021-05-23T19:53:31.819600Z","shell.execute_reply":"2021-05-23T19:53:32.513997Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torchvision import transforms","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:32.516802Z","iopub.execute_input":"2021-05-23T19:53:32.517131Z","iopub.status.idle":"2021-05-23T19:53:32.523310Z","shell.execute_reply.started":"2021-05-23T19:53:32.517094Z","shell.execute_reply":"2021-05-23T19:53:32.522302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_aug_train = df_train\ndf_preproc_val = df_val\ndf_preproc_test = df_test\nloader_transform = transforms.RandomRotation(180)","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:34.093221Z","iopub.execute_input":"2021-05-23T19:53:34.093589Z","iopub.status.idle":"2021-05-23T19:53:34.097833Z","shell.execute_reply.started":"2021-05-23T19:53:34.093557Z","shell.execute_reply":"2021-05-23T19:53:34.096992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pre_process(image, color=True, gaussian=False, kernel=IMG_SIZE//10):\n    if color:\n        image = image\n    else:\n        image = cv2.cvtColor(image, cv2.IMREAD_GRAYSCALE)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    image = cv2.resize(image, (IMG_SIZE, IMG_SIZE))\n    if gaussian:\n        image=cv2.addWeighted ( image,4, cv2.GaussianBlur( image , (0,0) , kernel) ,-4 ,128)\n    else:\n        image=cv2.addWeighted ( image,4, cv2.medianBlur(image, kernel) ,-4 ,128)\n        \n    return PIL.Image.fromarray(image, \"RGB\")","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:34.548809Z","iopub.execute_input":"2021-05-23T19:53:34.549120Z","iopub.status.idle":"2021-05-23T19:53:34.557763Z","shell.execute_reply.started":"2021-05-23T19:53:34.549090Z","shell.execute_reply":"2021-05-23T19:53:34.556880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef aug_train_img_creator(tranform, df_init=df_train, df_augmented=df_aug_train):\n    try:\n        shutil.rmtree('refuge_data')\n    except:\n        print(\"no such directory1\")\n    \n    shutil.copytree('/kaggle/input/eurecom-aml-2021-challenge-2/refuge_data/refuge_data','refuge_data/refuge_data' )\n    try:\n        shutil.rmtree('refuge_data/refuge_data/train/gts')\n        shutil.rmtree('refuge_data/refuge_data/train/images')\n        os.remove('refuge_data/refuge_data/train/index.json')\n    except:\n        print(\"no such directory2\")\n    \n    os.mkdir('refuge_data/refuge_data/train/images')\n    os.mkdir('refuge_data/refuge_data/train/gts')\n    i = 0\n    for image in df_init.IMG_NAME:\n        i = i+1\n        if i%10 ==0:\n            print(i)\n        path_to_img = dir_path+\"train/images/\"+image\n        bmp = image.replace('.jpg', '.bmp')\n        path_to_bmp = dir_path+\"train/gts/\"+bmp\n        img = cv2.imread(path_to_img)\n        bmpimg = PIL.Image.open(path_to_bmp)\n        #original = pre_process(img)\n        original = PIL.Image.fromarray(img, \"RGB\")\n        name = image.replace('.jpg', '')\n        bmpimg.save(\"refuge_data/refuge_data/train/gts/\"+bmp)\n        original.save(\"refuge_data/refuge_data/train/images/\"+name+\".jpg\")\n        for k in range(6):\n            new_img = loader_transform(original)\n            name = image.replace('.jpg', '') + str(k)\n            new_img.save(\"refuge_data/refuge_data/train/images/\"+name+\".jpg\")\n            bmpimg.save(\"refuge_data/refuge_data/train/gts/\"+name+\".bmp\")\n            new_sample = df_init[df_init.IMG_NAME == image]\n            new_sample.IMG_NAME = name+\".jpg\"\n            df_augmented = df_augmented.append(new_sample, ignore_index=True)\n    print('Augmentation ok')\n    return df_augmented","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:36.416265Z","iopub.execute_input":"2021-05-23T19:53:36.416650Z","iopub.status.idle":"2021-05-23T19:53:36.427132Z","shell.execute_reply.started":"2021-05-23T19:53:36.416621Z","shell.execute_reply":"2021-05-23T19:53:36.426069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_val","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:37.070248Z","iopub.execute_input":"2021-05-23T19:53:37.070601Z","iopub.status.idle":"2021-05-23T19:53:37.092696Z","shell.execute_reply.started":"2021-05-23T19:53:37.070570Z","shell.execute_reply":"2021-05-23T19:53:37.092009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_val_test( df_vinit = df_val , df_tinit = df_test):\n    try:\n        shutil.rmtree('refuge_data/refuge_data/val/images')\n        shutil.rmtree('refuge_data/refuge_data/val/gts')\n        os.remove('refuge_data/refuge_data/val/index.json')\n        shutil.rmtree('refuge_data/refuge_data/test/images')\n        os.remove('refuge_data/refuge_data/test/index.json')\n    except:\n        print(\"no such directory2\")\n    \n    os.mkdir('refuge_data/refuge_data/val/images')\n    os.mkdir('refuge_data/refuge_data/val/gts')\n    os.mkdir('refuge_data/refuge_data/test/images')\n    df_v = pd.DataFrame(columns=['IMG_NAME','Label', 'Fovea_X', 'Fovea_Y', 'Size_X', 'Size_Y'])\n    df_t = pd.DataFrame(columns=['IMG_NAME', 'Size_X', 'Size_Y'])\n    for image in df_vinit.IMG_NAME:\n        path_to_img = dir_path+\"val/images/\"+image\n        path_to_bmp = dir_path+\"val/gts/\"+image.replace('.jpg', '.bmp')\n        img = cv2.imread(path_to_img)\n        bmp = PIL.Image.open(path_to_bmp)\n        #new_img = pre_process(img)\n        new_img = PIL.Image.fromarray(img, \"RGB\")\n        new_img.save(\"refuge_data/refuge_data/val/images/\"+image)\n        bmp.save(\"refuge_data/refuge_data/val/gts/\"+image.replace('.jpg', '.bmp'))\n        sample = df_vinit[df_vinit.IMG_NAME == image]\n        df_v = df_v.append(sample, ignore_index=True)\n    \n    print('Preprocessing validation ok')\n    for image in df_tinit.IMG_NAME:\n        path_to_img = dir_path+\"test/images/\"+image\n        img = cv2.imread(path_to_img)\n        #new_img = pre_process(img)\n        new_img = PIL.Image.fromarray(img, \"RGB\")\n        new_img.save(\"refuge_data/refuge_data/test/images/\"+image)\n        sample = df_tinit[df_tinit.IMG_NAME == image]\n        df_t = df_t.append(sample, ignore_index=True)\n    print('Preprocessing test ok')\n    return df_v, df_t\n        ","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:38.189398Z","iopub.execute_input":"2021-05-23T19:53:38.189711Z","iopub.status.idle":"2021-05-23T19:53:38.199469Z","shell.execute_reply.started":"2021-05-23T19:53:38.189682Z","shell.execute_reply":"2021-05-23T19:53:38.198339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loader_transform = transforms.RandomRotation(10)\ndf_aug_train = aug_train_img_creator(loader_transform)","metadata":{"execution":{"iopub.status.busy":"2021-05-23T19:53:38.768153Z","iopub.execute_input":"2021-05-23T19:53:38.768504Z","iopub.status.idle":"2021-05-23T20:01:36.947837Z","shell.execute_reply.started":"2021-05-23T19:53:38.768471Z","shell.execute_reply":"2021-05-23T20:01:36.939784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_preproc_val, df_preproc_test = preprocess_val_test()","metadata":{"execution":{"iopub.status.busy":"2021-05-23T20:01:37.230309Z","iopub.execute_input":"2021-05-23T20:01:37.230673Z","iopub.status.idle":"2021-05-23T20:03:29.220335Z","shell.execute_reply.started":"2021-05-23T20:01:37.230635Z","shell.execute_reply":"2021-05-23T20:03:29.219392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def write_jsons(train=df_aug_train, val=df_preproc_val, test=df_preproc_test):\n    train = train.T\n    val = val.T\n    test = test.T\n    train.to_json('refuge_data/refuge_data/train/index.json')\n    val.to_json('refuge_data/refuge_data/val/index.json')\n    test.to_json('refuge_data/refuge_data/test/index.json')\n    \n    print('Write JSONs ok')","metadata":{"execution":{"iopub.status.busy":"2021-05-23T20:03:29.237531Z","iopub.execute_input":"2021-05-23T20:03:29.237889Z","iopub.status.idle":"2021-05-23T20:03:29.245972Z","shell.execute_reply.started":"2021-05-23T20:03:29.237850Z","shell.execute_reply":"2021-05-23T20:03:29.245262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"write_jsons()","metadata":{"execution":{"iopub.status.busy":"2021-05-23T20:03:29.247740Z","iopub.execute_input":"2021-05-23T20:03:29.248137Z","iopub.status.idle":"2021-05-23T20:03:29.413203Z","shell.execute_reply.started":"2021-05-23T20:03:29.248099Z","shell.execute_reply":"2021-05-23T20:03:29.412281Z"},"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":"import matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2021-05-23T20:03:29.414473Z","iopub.execute_input":"2021-05-23T20:03:29.414828Z","iopub.status.idle":"2021-05-23T20:03:29.419200Z","shell.execute_reply.started":"2021-05-23T20:03:29.414790Z","shell.execute_reply":"2021-05-23T20:03:29.418063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root_dir = 'refuge_data/refuge_data'\nlr = 1e-4\nbatch_size = 8\nnum_workers = 8\ntotal_epoch = 100\n# Datasets\n\n\ntrain_set = RefugeDataset(root_dir, \n                          split='train')\n\nval_set = RefugeDataset(root_dir, \n                        split='val')\n\ntest_set = RefugeDataset(root_dir, \n                         split='test')\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-23T20:03:29.420630Z","iopub.execute_input":"2021-05-23T20:03:29.421256Z","iopub.status.idle":"2021-05-23T20:17:00.748241Z","shell.execute_reply.started":"2021-05-23T20:03:29.421219Z","shell.execute_reply":"2021-05-23T20:17:00.746420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Device, model, loss and optimizer","metadata":{}},{"cell_type":"code","source":"import torchvision.models as models\n# Device\ndevice = torch.device(\"cuda:0\")\n\nmodel = UNet(n_channels=3, n_classes=2).to(device)\n# model = models.inception_v3(pretrained=True).to(device)\n# model.AuxLogits.fc = nn.Linear(768, 2)\n# model.fc = nn.Linear(2048, 2)\n\n#model = models.resnet50(pretrained=True).to(device)\n#model.fc=nn.Linear(512, 2)\n\n","metadata":{"execution":{"iopub.status.busy":"2021-05-23T20:30:38.213116Z","iopub.execute_input":"2021-05-23T20:30:38.215735Z","iopub.status.idle":"2021-05-23T20:30:38.595814Z","shell.execute_reply.started":"2021-05-23T20:30:38.215688Z","shell.execute_reply":"2021-05-23T20:30:38.594787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Loss\nseg_loss = torch.nn.BCELoss(reduction='mean')\n\n# Optimizer\noptimizer = optim.Adam(model.parameters(), lr=lr)","metadata":{},"execution_count":null,"outputs":[]},{"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.\nepoch = 50\nepoch_c=0\n\nwhile epoch_c < total_epoch:\n    epoch_c+=1\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        print(logits.shape)\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        \n","metadata":{"execution":{"iopub.status.busy":"2021-05-23T20:26:27.678765Z","iopub.execute_input":"2021-05-23T20:26:27.679086Z","iopub.status.idle":"2021-05-23T20:26:28.581873Z","shell.execute_reply.started":"2021-05-23T20:26:27.679055Z","shell.execute_reply":"2021-05-23T20:26:28.579701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load best model + classifier","metadata":{}},{"cell_type":"code","source":"# Load model and classifier\nmodel = UNet(n_channels=3, n_classes=2).to(device)\nmodel.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 = pickle.load(clf_file)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Check performance is maintained on validation","metadata":{}},{"cell_type":"code","source":"model.eval()\nval_vCDRs = []\nval_classif_gts = []\nval_loss = 0.\nval_dsc_od = 0.\nval_dsc_oc = 0.\nval_vCDR_error = 0.\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(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.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":"markdown","source":"# Predictions on test set","metadata":{}},{"cell_type":"code","source":"nb_test_batches = len(test_loader)\nmodel.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(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            \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.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)\n\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":{"trusted":true},"execution_count":null,"outputs":[]}]}