{"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":"# Import","metadata":{"id":"Fez-IK669ngB"}},{"cell_type":"code","source":"from google.colab import drive\ndrive.mount('/content/gdrive')","metadata":{"id":"lhQE-Psw9qA7","outputId":"44b9369d-d0f8-4809-e5fe-8e9f3f651d8e"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cd \"/content/gdrive/MyDrive/Dataset/\"","metadata":{"id":"mjl49ZEu9qTw"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport math\nimport random\nimport os\nfrom functools import partial\nimport seaborn as sn\n\nimport matplotlib.pyplot as plt\nfrom matplotlib.pyplot import figure\nimport matplotlib.gridspec as gridspec\n\nfrom PIL import Image , ImagePalette, ImageDraw, ImageFont\nfrom IPython.display import Image as img\nimport cv2\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import MultiLabelBinarizer\nfrom sklearn.preprocessing import MinMaxScaler\nfrom sklearn.metrics import roc_curve, auc\n\nimport torchvision\nfrom torchvision import transforms\nfrom torchvision import models\nfrom torchvision.utils import make_grid\nfrom torchvision import transforms, utils\nimport torchvision.utils as vutils\n\nimport torch\nfrom torch.autograd import Variable\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler","metadata":{"id":"MEh29_dI9ngH"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random.seed(42)","metadata":{"id":"p2Xy_e4dyPC6"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configurations","metadata":{"id":"_5IGlZdBeq0G"}},{"cell_type":"code","source":"LR = 0.001\nMOMENTUM = 0.9\nWEIGHT_DECAY = 0.00005\nPAMR_KERNEL_SIZE = 3\nPAMR_DILATIONS = [1,2,4,8,12,24]\nSG_ALPHA = 0.5\nFOCAL_PENALTY = 3\nFOCAL_LAMBDA = 0.01\n\n\nimage_dir = '/content/gdrive/MyDrive/Dataset/image/'\nCSV_PATH = \"/content/gdrive/MyDrive/Dataset/train.csv\"","metadata":{"id":"2H_6PUyYeuwd"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Analysis","metadata":{"id":"W_OFxIxF9ngH"}},{"cell_type":"code","source":"Label_number = ['0','1','10','11','12','13','14','15','16','17','18','2','3','4','5','6','7','8','9']\n\n\n\nClasses= ['Nucleoplasm',\n          'Nuclear membrane',\n          'Microtubules',\n          'Mitotic spindle',\n          'Centrosome',\n          'Plasma membrane',\n          'Mitochondria',\n          'Aggresome',\n          'Cytosol',\n          'Vesicles and punctate cytosolic patterns',\n          'Negative',\n          'Nucleoli',\n          'Nucleoli fibrillar center',\n          'Nuclear speckles',\n          'Nuclear bodies',\n          'Endoplasmic reticulum',\n          'Golgi apparatus',\n          'Intermediate filaments',\n          'Actin filaments'\n          ]\n\n\n# Colour 0,0,0 is an extra colour for visualising background\n\nColours = [\n           0,0,0,\n           128,128,128,\n           255,255,255,\n           240,50,230,\n           145,30,180,\n           220,195,255,\n           0,0,128,\n           0,130,200,\n           70,240,240,\n           0,128,128,\n           60,180,75,\n           170,255,195,\n           210,245,60,\n           128,128,0,\n           255,225,25,\n           255,250,200,\n           170,110,40,\n           245,130,48,\n           255,215,180,\n           128,0,0\n          ]\n\n","metadata":{"id":"UL2ddrPR9ngI"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataframe=pd.read_csv(CSV_PATH)\ndataframe.head()","metadata":{"id":"y8e5vSlG9ngI"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mlb = MultiLabelBinarizer()\nmlb.fit(dataframe.iloc[:,1].str.split('|'))\nmlb.classes_","metadata":{"id":"h1BPgVg79ngI"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizing images and classes","metadata":{"id":"3L_R2VKN9ngJ"}},{"cell_type":"code","source":"# Get image id that contains only a single label for each class\n\nimgIdOneLabel = pd.DataFrame(columns=dataframe.columns)\nimgIdOneLabel= []\nfor i in range(19):\n    imgIdOneLabel.append(dataframe[dataframe['Label'] == str(i)].iloc[0])\nimgIdOneLabel","metadata":{"id":"5in28xxr9ngJ"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# A function for retrieving image\ndef getImage(imgPath):\n    # Image downscaled to resolutions of 1440 to reduce training time\n    MAX_SIZE = (1440,1440)\n    greenPath = imgPath+'_green.png'\n\n    # Only grayscale image for green channel is used\n    green = cv2.imread(greenPath, cv2.IMREAD_GRAYSCALE)\n\n    # Green channel is padded with empty red and blue channel\n    red = np.zeros((green.shape[1], green.shape[0], 1), dtype = \"uint8\")\n    blue = np.zeros((green.shape[1], green.shape[0], 1), dtype = \"uint8\")\n    rgb = cv2.merge((blue,green,red))\n\n    rgb_pil = Image.fromarray(rgb)\n    rgb_pil = Image.fromarray(rgb)\n    rgb_pil.thumbnail(MAX_SIZE)\n    \n    return rgb_pil","metadata":{"id":"UvGbix6R9ngJ"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show all of the classes and ther correspinding colour\ndef showClassesAndColours(idList):\n    \n    for i in range(19):\n        fig = plt.figure(figsize=(20,40))\n\n        print(\"Class:\",Classes[i])\n        print(\"Label number:\",Label_number[i])\n\n        gs = gridspec.GridSpec(1,2, width_ratios=[3,3], height_ratios=[1])\n        ax1 = fig.add_subplot(gs[0])\n        ax2 = fig.add_subplot(gs[1])\n\n        tensor = torch.ones((720,720))\n        mask = tensor.new_full((720,720),i)\n        mask = Image.fromarray(mask.numpy().astype('uint8'))\n        mask.putpalette(Colours[3:])  # Colour[0:3] RGB colour set is for background, thus ignored\n        \n        imgId = idList[i][0]\n        imgPath =  image_dir + imgId\n        ax1.imshow(getImage(imgPath))\n        ax2.imshow(mask)\n        plt.show()","metadata":{"id":"pj4rcjdE7snF"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"showClassesAndColours(imgIdOneLabel)","metadata":{"id":"kAcF_OGa7BIw"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define custom dataset class and dataloader\n","metadata":{"id":"3vP0AhFk9ngK"}},{"cell_type":"code","source":"class MultiLabelCellImage(Dataset):\n    def __init__(self, dataframe, image_dir, transform = None):\n        self.dataframe = dataframe\n        self.image_dir = image_dir\n        self.transform = transform\n        self.labels = mlb.transform(dataframe.iloc[:,1].str.split('|')).astype(np.int64)\n\n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, index):\n\n        #Get image ID and labels\n        img_id = self.dataframe.iloc[index,0]\n        label = torch.from_numpy(self.labels[index])\n\n        #Load image\n        imgPath = os.path.join(self.image_dir, img_id)\n        image = getImage(imgPath)\n\n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label, imgPath","metadata":{"id":"isdbmHhT9ngK"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trf = transforms.Compose([\n transforms.RandomResizedCrop(720),\n transforms.RandomHorizontalFlip(),\n transforms.ToTensor(),\n transforms.Normalize((0,0.5,0), (1,0.5,1))\n])","metadata":{"id":"Lp8X4jrp9ngK"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, test_df = train_test_split(dataframe, test_size=0.2)\n\n# Training dataset will be defined, and further split into train and val set in function train()\ntest_dataset= MultiLabelCellImage(dataframe=test_df, image_dir = image_dir, transform=trf)\ntest_loader = DataLoader(test_dataset, batch_size=2,shuffle=True, num_workers=2)","metadata":{"id":"YZFdRDP_9ngL"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Single-Stage Semantic Segmentation from Image Labels","metadata":{"id":"JrSK_zzZ9ngL"}},{"cell_type":"markdown","source":"## Atrous Spatial Pyramid Pooling (ASPP)","metadata":{"id":"V1U1az7k9ngL"}},{"cell_type":"code","source":"\ndef create_ASPP_block(in_channels, out_channels, kernel_size, stride, padding, dilation):\n    \n    aspp_block = nn.Sequential(nn.Conv2d(in_channels,out_channels,kernel_size, stride, padding, dilation,bias=False),\n                              nn.BatchNorm2d(out_channels))\n    \n    return aspp_block","metadata":{"id":"3jpHeLP49ngM"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Configuration retrieved from author with modification\nclass ASPP(nn.Module):\n    def __init__(self,in_channels,out_channels, kernel_size, stride, padding, dilation):\n        super().__init__()\n\n        self.number_of_blocks = len(kernel_size)\n        \n        self.aspp_blocks = nn.ModuleList([create_ASPP_block(in_channels,out_channels,kernel_size[i], stride[i], padding[i], dilation[i]) for i in range(self.number_of_blocks)])\n        \n        for i in self.modules():\n            if isinstance(i,nn.Conv2d):\n                torch.nn.init.kaiming_normal_(i.weight)\n                \n        # Retrieved from author\n        self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),\n                                     nn.Conv2d(in_channels, out_channels, 1, stride=1, bias=False),\n                                     nn.BatchNorm2d(out_channels),\n                                     nn.ReLU())\n        \n        self.conv1 = nn.Conv2d(out_channels*self.number_of_blocks,out_channels,1,bias=False).cuda()\n        self.bn =  nn.BatchNorm2d(out_channels)\n        self.dropout = nn.Dropout(0.5)\n    def forward(self,x):\n\n        x_arr = [None] * (self.number_of_blocks+1)\n        for i in range(self.number_of_blocks):\n            curBlock = self.aspp_blocks[i]\n            temp = curBlock(x)\n            x_arr[i] = F.relu(temp)\n            \n            \n        temp = self.global_avg_pool(x)\n        \n\n        x_arr[self.number_of_blocks] = F.interpolate(temp, size=x_arr[self.number_of_blocks - 1].size()[2:], mode='bilinear', align_corners=True)\n        \n\n        if x.shape[0]>=2:\n            temp = torch.cat((x_arr[0],x_arr[1]),dim=1)\n            for i in range(2,self.number_of_blocks):\n                temp = torch.cat((temp,x_arr[i]),dim = 1)\n            \n        x = temp\n        \n        x = self.conv1(x)\n        x = self.bn(x)\n        x = F.relu(x)\n        x = self.dropout(x)\n        \n        return x","metadata":{"id":"JOrV8rAvx9RG"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Global Cue Injection (GCI)","metadata":{"id":"QPlPvdCH9ngO"}},{"cell_type":"code","source":"# Configuration retrieved from author\nclass GCI(nn.Module):\n    def __init__(self,in_channels):\n        super().__init__()\n        self.conv_deep = nn.Sequential(nn.Conv2d(in_channels,in_channels*2,1,bias=False),\n                                      nn.BatchNorm2d(in_channels*2))\n        self.conv_shallow = nn.Sequential(nn.Conv2d(in_channels, in_channels,1,bias=False),\n                                         nn.BatchNorm2d(in_channels))\n        self.conv_cls = nn.Sequential(nn.Conv2d(in_channels,in_channels,1,bias=False),\n                                     nn.BatchNorm2d(in_channels))\n        for i in self.modules():\n            if isinstance(i,nn.Conv2d):\n                torch.nn.init.kaiming_normal_(i.weight)\n                \n    def adin(self,x_shallow,x_deep):\n\n        assert x_deep.shape[1]%2==0\n        half_size = int(x_deep.shape[1]/2)\n\n        # Split the max_vec from x_deep to half\n        z = x_deep[:,:half_size]\n        b = x_deep[:,half_size:]\n\n        # Unsqueeze z and b to be same dimensions as x_shallow\n        z = z.unsqueeze(-1).unsqueeze(-1)\n        b = b.unsqueeze(-1).unsqueeze(-1)\n        return F.relu((z*(x_shallow)+1) + b)\n    \n    \n    def forward(self,x_shallow,x_deep):\n\n        x_deep = F.relu(self.conv_deep(x_deep))\n        x_shallow = self.conv_shallow(x_shallow)\n\n        # Flatten the receptive field and perform global max pooling on it\n        max_vec,indexes = x_deep.view(x_deep.shape[0],x_deep.shape[1],-1).max(-1)\n\n        x = self.adin(x_shallow,max_vec)\n        x = self.conv_cls(x)\n        \n        return x","metadata":{"id":"P-9S4skI9ngO"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Stochastic Gate (SG)","metadata":{"id":"6olaAFyB9ngP"}},{"cell_type":"code","source":"# COnfiguration retrieved from author\ndef StochasticGate(x_shallow,x_deep,a=SG_ALPHA,mode='Training'):\n    b = 1/(1-a)\n    \n    if mode == 'Training':\n\n        # Get a random sequence of 0 and 1 based on SG_ALPHA\n        rand = torch.rand(x_deep.shape[1])\n        rand = torch.where(rand > a, 1, 0)\n        if torch.cuda.is_available():\n          rand = rand.cuda()\n\n        # Randomly mixing x_deep and x_shallow based on the random sequence\n        x_sg = ((1-rand)*b*(x_deep.permute(0,2,3,1)-a*x_shallow.permute(0,2,3,1))+rand*x_shallow.permute(0,2,3,1)).permute(0,3,1,2)\n        return x_sg\n\n        # Evaluating mode will mix x_deep and x_shallow deterministically\n    elif mode == 'Evaluating':\n        x_sg = (1-a)*x_deep+a*x_shallow\n        \n        return x_sg\n    else:\n        print('Unknown mode, mode should be Training or Evaluating')","metadata":{"id":"We5_KFN89ngP"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pixel-Adaptive Mask Refinement (PAMR)","metadata":{"id":"KQw3GS9T9ngP"}},{"cell_type":"code","source":"# Retrieved from author with partial modification\nclass LocalAffinity(nn.Module):\n\n    def __init__(self, kernel_size=3, dilations=[1]):\n        super(LocalAffinity, self).__init__()\n        assert kernel_size%2 == 1 and kernel_size>=3\n        \n        self.kernel_size=kernel_size\n        self.dilations = dilations\n        weight = self._init_aff()\n        self.register_buffer('kernel', weight)\n\n    def _init_aff(self):\n        # initialising the shift kernel\n        \n        weight = torch.zeros(self.kernel_size**2-1,1,self.kernel_size,self.kernel_size)\n        center_pos = math.floor(self.kernel_size/2)\n        for i in range(weight.size(0)):\n            weight[i, 0, center_pos, center_pos] = 1\n\n\n        cur_h = 0\n        cur_w = 0\n        after_center = 0\n        for i in range(weight.size(0)):\n                    if cur_h==center_pos and cur_w==center_pos:\n                        after_center=1\n\n                    weight[i,0,cur_h,cur_w+after_center] = -1\n\n                    cur_w+=1\n                    if (cur_w+after_center) == self.kernel_size:\n                        cur_w=0\n                        cur_h+=1\n                        after_center=0\n                    \n        self.weight_check = weight.clone()\n\n        return weight\n\n    def forward(self, x):\n        \n        self.weight_check = self.weight_check.type_as(x)\n        assert torch.all(self.weight_check.eq(self.kernel))\n\n        B,K,H,W = x.size()\n        x = x.view(B*K,1,H,W)\n        k = self.kernel_size\n        x_affs = []\n        for d in self.dilations:\n            x_pad = F.pad(x, [ int(((k+(d-1)*(k-1))-1)/2) ]*4, mode='replicate')\n            x_aff = F.conv2d(x_pad, self.kernel, dilation=d)\n            x_affs.append(x_aff)\n\n        x_aff = torch.cat(x_affs, 1)\n        return x_aff.view(B,K,-1,H,W)\n  \n\n\nclass LocalAffinityCopy(LocalAffinity):\n\n    def _init_aff(self):\n        # initialising the shift kernel\n        \n        weight = torch.zeros(self.kernel_size**2-1,1,self.kernel_size,self.kernel_size)\n        center_pos = math.floor(self.kernel_size/2)\n\n        cur_h = 0\n        cur_w = 0\n        after_center = 0\n        for i in range(weight.size(0)):\n                    if cur_h==center_pos and cur_w==center_pos:\n                        after_center=1\n\n                    weight[i,0,cur_h,cur_w+after_center] = 1\n\n                    cur_w+=1\n                    if (cur_w+after_center) == self.kernel_size:\n                        cur_w=0\n                        cur_h+=1\n                        after_center=0\n\n                    \n        self.weight_check = weight.clone()\n        return weight\n\n    \nclass LocalStDev(LocalAffinity):\n    \n    def _init_aff(self):\n\n        \n        weight = torch.zeros(self.kernel_size**2,1,self.kernel_size,self.kernel_size)\n        center_pos = math.floor(self.kernel_size/2)\n\n        cur_h = 0\n        cur_w = 0\n\n        for i in range(weight.size(0)):\n                    weight[i,0,cur_h,cur_w] = 1\n\n                    cur_w += 1\n                    if cur_w == self.kernel_size:\n                        cur_w=0\n                        cur_h+=1\n\n        self.weight_check = weight.clone()\n        return weight\n\n    def forward(self, x):\n        # returns (B,K,P,H,W), where P is the number\n        # of locations\n        x = super(LocalStDev, self).forward(x)\n\n        return x.std(2, keepdim=True)\n\n    \nclass LocalAffinityAbs(LocalAffinity):\n\n    def forward(self, x):\n        x = super(LocalAffinityAbs, self).forward(x)\n        return torch.abs(x)\n    \n    \nclass PAMR(nn.Module):\n\n    def __init__(self, num_iter=1, kernel_size=3,dilations=[1]):\n        super(PAMR, self).__init__()\n\n        self.num_iter = num_iter\n        self.aff_x = LocalAffinityAbs(kernel_size,dilations)\n        self.aff_m = LocalAffinityCopy(kernel_size,dilations)\n        self.aff_std = LocalStDev(kernel_size,dilations)\n\n    def forward(self, x, mask):\n        mask = F.interpolate(mask, size=x.size()[-2:], mode=\"bilinear\", align_corners=True)\n\n        # x: [BxKxHxW]\n        # mask: [BxCxHxW]\n        B,K,H,W = x.size()\n        _,C,_,_ = mask.size()\n\n        x_std = self.aff_std(x)\n\n        x = -self.aff_x(x) / (1e-8 + 0.1 * x_std)\n        x = x.mean(1, keepdim=True)\n        x = F.softmax(x, 2)\n\n        for _ in range(self.num_iter):\n            m = self.aff_m(mask)  # [BxCxPxHxW]\n            mask = (m * x).sum(2)\n\n        # xvals: [BxCxHxW]\n        return mask","metadata":{"id":"2Crv_OHidpdR"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Network","metadata":{"id":"ARXaW8FO9ngR"}},{"cell_type":"markdown","source":"## Helper function","metadata":{"id":"q_kk-ytJ9ngR"}},{"cell_type":"code","source":"# Retrieved from author\n\ndef rescale_as(x, y, mode=\"bilinear\", align_corners=True):\n    h, w = y.size()[2:]\n    x = F.interpolate(x, size=[h, w], mode=mode, align_corners=align_corners)\n    return x\n\ndef pseudo_gtmask(mask, cutoff_top=0.6, cutoff_low=0.2, eps=1e-8):\n    \"\"\"Convert continuous mask into binary mask\"\"\"\n    bs,c,h,w = mask.size()\n    mask = mask.view(bs,c,-1)\n\n    # for each class extract the max confidence\n    mask_max, _ = mask.max(-1, keepdim=True)\n    mask_max[:, :1] *= 0.7\n    mask_max[:, 1:] *= cutoff_top\n    #mask_max *= cutoff_top\n\n    # if the top score is too low, ignore it\n    lowest = torch.Tensor([cutoff_low]).type_as(mask_max)\n    mask_max = mask_max.max(lowest)\n\n    pseudo_gt = (mask > mask_max).type_as(mask)\n\n    # remove ambiguous pixels\n    ambiguous = (pseudo_gt.sum(1, keepdim=True) > 1).type_as(mask)\n    pseudo_gt = (1 - ambiguous) * pseudo_gt\n\n    return pseudo_gt.view(bs,c,h,w)\n\ndef balanced_mask_loss_ce(mask, pseudo_gt, gt_labels, ignore_index=255):\n    \"\"\"Class-balanced CE loss\n    - cancel loss if only one class in pseudo_gt\n    - weight loss equally between classes\n    \"\"\"\n\n    mask = F.interpolate(mask, size=pseudo_gt.size()[-2:], mode=\"bilinear\", align_corners=True)\n    \n    # indices of the max classes\n    mask_gt = torch.argmax(pseudo_gt, 1)\n\n    # for each pixel there should be at least one 1\n    # otherwise, ignore\n    ignore_mask = pseudo_gt.sum(1) < 1.\n    mask_gt[ignore_mask] = ignore_index\n\n    # class weight balances the loss w.r.t. number of pixels\n    # because we are equally interested in all classes\n    bs,c,h,w = pseudo_gt.size()\n    num_pixels_per_class = pseudo_gt.view(bs,c,-1).sum(-1)\n    num_pixels_total = num_pixels_per_class.sum(-1, keepdim=True)\n    class_weight = (num_pixels_total - num_pixels_per_class) / (1 + num_pixels_total)\n    class_weight = (pseudo_gt * class_weight[:,:,None,None]).sum(1).view(bs, -1)\n\n    # BCE loss\n    loss = F.cross_entropy(mask, mask_gt, ignore_index=ignore_index, reduction=\"none\")#,weight = weight\n    loss = loss.view(bs, -1)\n\n    # we will have the loss only for batch indices\n    # which have all classes in pseudo mask\n    gt_num_labels = gt_labels.sum(-1).type_as(loss) + 1 # + BG\n    ps_num_labels = (num_pixels_per_class > 0).type_as(loss).sum(-1)\n    batch_weight = (gt_num_labels == ps_num_labels).type_as(loss)\n\n    loss = batch_weight * (class_weight * loss).mean(-1)\n    return loss\n\ndef _rescale_and_clean(masks, image, labels):\n    \"\"\"Rescale to fit the image size and remove any masks\n    of labels that are not present\"\"\"\n    masks = F.interpolate(masks, size=image.size()[-2:], mode='bilinear', align_corners=True)\n    masks[:, 1:] *= labels[:, :, None, None].type_as(masks)\n    return masks\n\nclass Flatten(nn.Module):\n    def forward(self, input):\n        return input.view(input.size(0), -1)","metadata":{"id":"ejNpLQkX9ngR"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Main network","metadata":{"id":"PC-aEIep79Hk"}},{"cell_type":"code","source":"# Configuration retrieved from author\nclass SingleStageWSEG(nn.Module):\n    def __init__(self,num_of_classes):\n        super().__init__()\n        self.num_of_classes = num_of_classes\n        \n        backbone = torchvision.models.resnet50(pretrained=True)\n        #Split the backbone to two parts to get shallow features\n        backbone1 = nn.Sequential(*list(backbone.children())[0:5])\n        backbone2 = nn.Sequential(*list(backbone.children())[5:-2])\n        \n\n        for name, param in backbone2.named_parameters():\n             if not any(name.startswith(ext) for ext in ['layer4']):\n                 param.requires_grad = False\n\n        \n        del backbone\n        self.backbone1 = backbone1\n        self.backbone2 = backbone2\n        \n\n        # Load pretrain\n        # Remove last layer,Freeze layers\n        self._init_aspp()\n        self._init_gci()\n        self._init_pamr()\n        \n        #Configuration retrieved from author\n        self.shallow_skip_conv = nn.Sequential(nn.Conv2d(256,48,1,bias=False),\n                                 nn.BatchNorm2d(48))\n        self.deep_feature_conv = nn.Sequential(nn.Conv2d(304,256,3,1,1,bias=False),\n                                 nn.BatchNorm2d(256))\n        self.conv1 = nn.Sequential(nn.Conv2d(256,256,3,1,bias=False),\n                                 nn.BatchNorm2d(256))\n        self.conv2 = nn.Sequential(nn.Conv2d(256,256,3,1,bias=False),\n                                 nn.BatchNorm2d(256))\n        self.conv_last = nn.Conv2d(256,self.num_of_classes,1,1)\n        \n    def _init_aspp(self):\n        kernel_size = [1,3,3,3]\n        stride = [1,1,1,1]\n        padding = [0,6,12,18]\n        dilation = [1,6,12,18]\n        \n        self.aspp = ASPP(2048,256,kernel_size,stride,padding,dilation)\n\n    def _init_gci(self):\n        self.gci = GCI(256)#.cuda()\n        \n    def _init_pamr(self):\n        \n        pamr_iter = 10\n        pamr_kernel_size = PAMR_KERNEL_SIZE\n        pamr_dilations = PAMR_DILATIONS\n        \n        self.pamr = PAMR(pamr_iter,pamr_kernel_size,pamr_dilations)\n        \n    \n    def forward(self,y,y_raw=None,labels=None):\n        if y_raw is None and labels is None:\n            mode = 'Evaluating'\n        else:\n            mode = 'Training'\n\n        # Forward y through backbone1 to get shallow features x1\n        x1 = self.backbone1(y)\n\n        # Forward x1 through backbone2 to get deep features\n        x = self.backbone2(x1)\n        x = self.aspp(x)\n        \n        \n        x_shallow = F.relu(self.shallow_skip_conv(x1))\n\n        #Resize deep features to be same size as shallow deatures and concat with shallow features\n        x_deep = F.interpolate(x,x_shallow.size()[2:])\n        x_deep = torch.cat([x_deep,x_shallow],1)\n        x_deep = F.relu(self.deep_feature_conv(x_deep))\n        \n        x_shallow = self.gci(x1,x_deep)\n        \n\n        x = StochasticGate(x_shallow,x_deep,SG_ALPHA,mode)\n        x = F.dropout(F.relu(self.conv1(x)),0.5)\n        x = F.dropout(F.relu(self.conv2(x)),0.1)\n        x = self.conv_last(x)    \n        # x is the segmentation predicted by the network\n\n        # Add a non-zero background channels to x before performing pixelwise softmax\n        # To prevent zero-division error\n        bg = torch.ones_like(x[:, :1])\n        x = torch.cat([bg, x], 1)\n            \n        # Perform pixelwise-softmax to obtain a mask with dimensions bs , num_of_classes + 1, h, w\n        # Each pixel in each channel represents the probability of the pixel to be the class\n        masks = F.softmax(x,dim=1)\n        bs, c, h, w = x.size()\n        \n        #Flatten predicted segmentation x, and the mask to calculate normalised global weighted pooling (ngwp)\n        features = x.view(bs, c, -1)\n        masks_ = masks.view(bs, c, -1)\n\n        # +1 in denominator to prevent zero division error\n        ngwp = (features * masks_).sum(-1) / (1.0 + masks_.sum(-1))\n\n        # Calculate the average value of each pixels for a particular channel (class)\n        masks_mean = masks_.mean(-1)\n\n        # Calculate focal_loss based on the mean of the masks\n        focal_lambda = FOCAL_LAMBDA\n        focal_penalty = FOCAL_PENALTY\n        focal_loss = ((1- masks_mean)**focal_penalty) * torch.log(focal_lambda + masks_mean)\n \n        # Adding the losses together to claculate the classification score\n        # The first channel is ignored as it is the background channel added\n        cls_score = ngwp[:, 1:] + focal_loss[:, 1:]\n \n        # If it is in evaluating mode, return the classification score and the mask predicted,else proceed to refine masks\n        if mode=='Evaluating':\n            return cls_score, F.interpolate(masks,y.size()[2:])\n\n        self._mask_logits = x\n\n        # Retrieved from author\n        # Foreground stats\n        masks_ = masks_[:, 1:]\n        cls_fg = (masks_.mean(-1) * labels).sum(-1) / labels.sum(-1)\n\n        # Mask refinement with PAMR\n        y_raw_ = F.interpolate(y_raw,masks.detach().size()[-2:], mode=\"bilinear\", align_corners=True)\n        masks_dec = self.pamr(y_raw_, masks.detach())\n\n        # Upscale the masks & clean\n        masks = _rescale_and_clean(masks, y, labels)\n        masks_dec = _rescale_and_clean(masks_dec, y, labels)\n\n        # Create pseudo GT\n        pseudo_gt = pseudo_gtmask(masks_dec).detach()\n        loss_mask = balanced_mask_loss_ce(self._mask_logits, pseudo_gt, labels)\n\n        return cls_score, cls_fg, {\"cam\": masks, \"dec\": masks_dec}, self._mask_logits, pseudo_gt, loss_mask","metadata":{"id":"cRIJ0SU59ngR"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"id":"seuSqYitdpdU"}},{"cell_type":"markdown","source":"## Creating network and loading checkpoint","metadata":{"id":"EEpB4OCP_Nw6"}},{"cell_type":"code","source":"net = SingleStageWSEG(19)\nnet = net.cuda()\noptimizer = optim.SGD(net.parameters(), lr = LR ,weight_decay=  WEIGHT_DECAY, momentum = MOMENTUM)","metadata":{"id":"CHVxVpkG9ngR"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ckpt = torch.load('/content/gdrive/MyDrive/Dataset/Model1pth/model1_106_.pth')","metadata":{"id":"NDPV42GK7L_A"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net.load_state_dict(ckpt['model_state_dict'])\nnet.backbone1.load_state_dict(ckpt['backbone1_state_dict'])\nnet.backbone2.load_state_dict(ckpt['backbone2_state_dict'])\nnet.aspp.load_state_dict(ckpt['aspp_state_dict'])\nnet.gci.load_state_dict(ckpt['gci_state_dict'])\nnet.pamr.load_state_dict(ckpt['pamr_state_dict'])\noptimizer.load_state_dict(ckpt['optimizer_state_dict'])","metadata":{"id":"f2iONMAegPnP"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Defining functions for training and validating","metadata":{"id":"uyibF1pO_RTz"}},{"cell_type":"code","source":"def train(net, optimizer, start_epochs ,end_epochs, start_i, train_df, verbose=True):\n \n train_loss_history={'loss_cls':0,\n                    'loss_fg':0,\n                    'loss_mask':0,\n                    'loss':0}\n\n # transfer model to GPU\n if torch.cuda.is_available():\n     net = net.cuda()\n\n \n train_loss=0\n # set to training mode\n net.train()\n # train the network\n for e in range(start_epochs,end_epochs): \n     running_loss = 0.0\n     running_count = 0.0\n\n     # Resplit the training set and validation set with ratio of 3:1 at every epoch\n     train_df, val_df = train_test_split(train_df, test_size=0.25)\n     train_dataset = MultiLabelCellImage(dataframe=train_df, image_dir = image_dir, transform=trf)\n     val_dataset = MultiLabelCellImage(dataframe=val_df, image_dir = image_dir, transform=trf)\n     train_loader = DataLoader(train_dataset, batch_size=4,shuffle=True, num_workers=4)\n     val_loader = DataLoader(val_dataset, batch_size=2,shuffle=True, num_workers=2)\n     loss_iterations = int(np.ceil(len(train_loader)/5))\n\n     for i, (images, labels, img_path) in enumerate(train_loader):\n         if e==start_epochs:\n             i = i+start_i\n         if e==start_epochs and i==len(train_loader):\n             break\n         print('Train:',i)\n         #Forward\n         # Denormalisation to get raw image values\n         images_raw = torch.zeros(images.shape)\n         for c, m, s in zip((0,1,2), (0,0.5,0), (1,0.5,1)):\n             images_raw[:,c,:,:] = images[:, c, :, :].mul_(s).add_(m)\n         \n         if torch.cuda.is_available():\n             images = images.cuda()\n             labels = labels.cuda()\n             images_raw = images_raw.cuda()\n\n         cls_score, cls_fg, masks, mask_logits, pseudo_gt, loss_mask = net(images,images_raw,labels)\n\n         #Masks\n         masks[\"pseudo\"] = pseudo_gt\n         for mask_key, mask_val in masks.items():\n             masks[mask_key] = masks[mask_key].detach()\n        \n         mask_logits = mask_logits.detach()\n\n         #Give weight of 3 to positive labels due to class imbalance problem\n         weight = labels*2+1\n         weight = weight.float()\n\n         loss_cls =  F.multilabel_soft_margin_loss(cls_score, labels.type(torch.DoubleTensor).cuda(), weight=weight, size_average=None).mean(0)\n         loss_fg = cls_fg.mean().item()\n         \n         if 'dec' in masks:\n             loss_mask = loss_mask.mean()\n         \n         loss = loss_cls.clone()\n\n         # Eaable the loss of mask after the 9th epochs\n         if e>=9:\n             loss += loss_mask\n\n         #Backpropagation\n         optimizer.zero_grad()\n         loss.backward()\n         optimizer.step()\n\n         # Put losses into a dictionary\n         train_loss_history['loss_cls']+=loss_cls\n         train_loss_history['loss_fg']+=loss_fg\n         train_loss_history['loss_mask']+=loss_mask.item()\n         train_loss_history['loss']+=loss\n            \n         # Calculate training loss\n         if i % loss_iterations == loss_iterations-1 or i == len(train_loader) - 1 or i==1:\n            for key in train_loss_history:\n                if i ==  len(train_loader):\n                    train_loss_history[key]=train_loss_history[key]/((i-start_i)%loss_iterations)\n                else:\n                    train_loss_history[key]=train_loss_history[key]/loss_iterations\n\n            # Visualise sample images and mask\n            visualiseSample(img_path,labels,masks,images_raw)    \n            net.eval()\n\n            #Perform validation\n            with torch.no_grad():\n                if i ==  len(train_loader):\n                    val_loss_history = validation(net,val_loader,(i%loss_iterations))\n                else:\n                    val_loss_history = validation(net,val_loader,loss_iterations)\n                \n                \n            net.train()\n            # Saving checkpoint\n            logging(e,i,net,train_loss_history,val_loss_history, verbose=True)\n            train_loss_history={'loss_cls':0,\n                                'loss_fg':0,\n                                'loss_mask':0,\n                                'loss':0\n            }\n     start_i=0","metadata":{"id":"VA_NDG5f9ngT"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validation(net,val_loader,n):\n    loss_iterations = int(np.ceil(len(val_loader)/5))\n    val_loss_history={'loss_cls':0}\n    for i, (images, labels, img_path) in enumerate(val_loader):\n         if i==n:\n             for key in val_loss_history:\n                    val_loss_history[key]=val_loss_history[key]/n\n             return val_loss_history\n        \n         print('Val:',i)\n         #Forward\n\n         images_raw = torch.zeros(images.shape)\n         for c, m, s in zip((0,1,2), (0,0.5,0.5), (1,0.5,1)):\n             images_raw[:,c,:,:] = images[:, c, :, :].mul_(s).add_(m)\n\n         if torch.cuda.is_available():\n             images = images.cuda()\n             labels = labels.cuda()\n             images_raw = images_raw.cuda()\n\n         cls_score, masks = net(images)\n\n         #Masks\n\n\n         weight = labels*2+1\n         weight = weight.float()\n         loss_cls =  F.multilabel_soft_margin_loss(cls_score,labels.type(torch.DoubleTensor).cuda(), weight=weight, size_average=None).mean(0)\n\n\n         loss = loss_cls.clone()\n         \n         val_loss_history['loss_cls']+=loss_cls\n         if i % loss_iterations == loss_iterations-1 or i == len(val_loader) - 1 or i==1:\n             for i in range(labels.shape[0]):\n              print('Labels: ',labels[i])\n    \n              print('Image:')\n              display(getImage(img_path[i]))\n    \n              print('Mask:')\n              mask = visualiseMasks(masks[i],images_raw[i])\n              display(mask)\n\n    return val_loss_history","metadata":{"id":"fhxEc7azdpdW"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def logging(e,i,net,train_losses,val_losses, verbose=True):\n\n\n    checkpoint_path = \"/content/gdrive/MyDrive/Dataset/Model1pth/E_\" + str(e) +\"_i\" + str(i) +\".pth\"\n    print('Train',train_losses)\n    print('Val',val_losses)\n\n    torch.save({\n    'epoch': e,\n    'model_state_dict':net.state_dict(),\n    'backbone1_state_dict': net.backbone1.state_dict(),\n    'backbone2_state_dict': net.backbone2.state_dict(),\n    'aspp_state_dict':net.aspp.state_dict(),\n    'gci_state_dict':net.gci.state_dict(),\n    'pamr_state_dict':net.pamr.state_dict(),\n    'optimizer_state_dict': optimizer.state_dict(),\n    'train_losses': train_losses,\n    'val_losses':val_losses    \n    }, checkpoint_path)\n    print('Checkpoint for epoch ',e,' and iterations ',i,' has been saved')\n","metadata":{"id":"DLe3B3Bu9ngT"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Visualise single mask\ndef visualiseMasks(mask,img_raw):\n\n    image = torch.argmax(mask, dim=0)\n\n    # Convert the denormalsed image tensor to numpy array\n    img_raw = (img_raw[1]*255).cpu().numpy().astype('uint8')\n\n    # Obtain nonzero value pixels in the image by dividing img_raw by itself, and fill nan with 0\n    bright_pixel=img_raw/img_raw\n    bright_pixel = np.nan_to_num(bright_pixel)\n\n    # Clean the mask based on the nonzero pixels in the image\n    image = bright_pixel*image.cpu().numpy().astype('uint8')\n\n    image = Image.fromarray(image.astype('uint8'))\n    \n    image.putpalette(Colours)\n\n    print(image)\n    return image","metadata":{"id":"hlAK7VrF9ngT"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualiseSample(img_path,labels,masks,img_raw):\n    for i in range(labels.shape[0]):\n        print('Labels: ',labels[i])\n        \n        print('Image:')\n        display(getImage(img_path[i]))\n        \n        print('Mask:')\n        maskCam = visualiseMasks(masks['cam'][i],img_raw[i])\n        display(maskCam)\n        \n        print('Refined mask by PAMR:')\n        maskDec = visualiseMasks(masks['dec'][i],img_raw[i])\n        display(maskDec)\n\n        print('Pseudo groundtruth:')\n        maskPseudo = visualiseMasks(masks['pseudo'][i],img_raw[i])\n        display(maskPseudo)","metadata":{"id":"SSEdQiiN9ngU"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{"id":"pbbMVg7x_fSk"}},{"cell_type":"code","source":"history,losses_history = train(net,optimizer,0,5,0,train_df)","metadata":{"id":"lkgY59SJBHgD"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluation and Analysis","metadata":{"id":"24T5qFRp8RrI"}},{"cell_type":"markdown","source":"## Training loss and Validation loss","metadata":{"id":"_7K0PY93AenD"}},{"cell_type":"code","source":"def show_losses(directory):\n\n  train_losses= []\n  val_losses=[] \n  steps =[]\n\n  i=0\n  for filename in os.listdir(directory):\n      if filename.endswith(\".pth\"):\n          filePath =  directory+filename\n          ckpt = torch.load(filePath)\n          train_losses.append( ckpt['train_losses']) \n          val_losses.append(ckpt['val_losses'])\n          steps.append(i)\n          i+=1\n          print(filename)\n\n\n  for key in train_losses[0]:\n      if key != 'loss_cls':\n          value = []\n          for i in range(len(train_losses)):\n              value.append(train_losses[i][key])\n\n          plt.plot(steps, value, label = 'Train loss for '+key)\n          plt.xlabel('Steps')\n          plt.ylabel(key)\n          plt.legend()\n          plt.show()\n\n  train_value = []\n  val_value = []\n  for i in range(len(train_losses)):\n      train_value.append(train_losses[i]['loss_cls'])\n      val_value.append(val_losses[i]['loss_cls'])\n  plt.plot(steps, train_value, label = 'Train loss for Classification loss')\n  plt.plot(steps, val_value, label = 'Validation loss for Classification loss')\n  plt.xlabel('Step')\n  plt.ylabel('Classification loss')\n  plt.legend()\n  plt.show()\n  return train_losses, val_losses, steps\n","metadata":{"id":"DIMt8z4g0p1s"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_losses, val_losses, steps = show_losses('/content/gdrive/MyDrive/Dataset/Model1pth/')","metadata":{"id":"otbH63xochMI"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Measuring performance with test set","metadata":{"id":"0Q-1P9WsA55a"}},{"cell_type":"code","source":"# Evaluate models based on test set  \ndef evaluation(net,test_loader):\n    net.eval()\n    loss_iterations = int(np.ceil(len(test_loader)/5))\n    test_loss = 0\n    confusion_matrix = torch.zeros((19,2,2)).cuda()\n    labels_arr = np.random.rand(1,19)\n    scores_arr = np.random.rand(1,19)\n\n    for i, (images, labels, img_path) in enumerate(test_loader):\n         print('Test:',i)\n         #Forward\n         images_raw = torch.zeros(images.shape)\n         for c, m, s in zip((0,1,2), (0,0.5,0.5), (1,0.5,1)):\n             images_raw[:,c,:,:] = images[:, c, :, :].mul_(s).add_(m)\n\n         if torch.cuda.is_available():\n             images = images.cuda()\n             labels = labels.cuda()\n             images_raw = images_raw.cuda()\n\n         cls_score, masks = net(images)\n         \n         #Masks\n         \n         sig_score = torch.sigmoid(cls_score)\n         labels_arr=np.append(labels_arr,labels.cpu().numpy(),axis=0)\n         scores_arr=np.append(scores_arr,sig_score.cpu().numpy(),axis=0)\n         pred_labels = (sig_score>0.5).int()\n         confusion_matrix = updateConfusionMatrix(labels, pred_labels, confusion_matrix,19)\n         weight = labels*2+1\n         weight = weight.float()\n         loss_cls =  F.multilabel_soft_margin_loss(cls_score,labels.type(torch.DoubleTensor).cuda(), weight=weight, size_average=None).mean(0)\n\n\n         loss = loss_cls.clone()\n         \n         test_loss += loss\n         if i % loss_iterations == loss_iterations-1 or i == len(test_loader) - 1 or i==1:\n             for i in range(labels.shape[0]):\n              print('Labels: ',labels[i])\n    \n              print('Image:')\n              display(getImage(img_path[i]))\n    \n              print('Mask:')\n              mask = visualiseMasks(masks[i],images_raw[i])\n              display(mask)\n            \n\n    test_loss=test_loss/len(test_loader)\n    \n\n    return test_loss, confusion_matrix, labels_arr[1:,:], scores_arr[1:,:]","metadata":{"id":"-JurMpljcpIF"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Update a given confusion matri based on labels and predicted labels\ndef updateConfusionMatrix(labels,pred_labels,confusion_matrix,num_of_classes):\n\n  precision=0\n  recall=0\n  for i in range(num_of_classes):\n\n    tp=torch.sum((pred_labels[:,i]==1) & (labels[:,i]==1))\n    confusion_matrix[i,0,0] += tp\n    \n    fp=torch.sum((pred_labels[:,i]==1) & (labels[:,i]==0))\n    confusion_matrix[i,0,1] += fp\n    \n    fn=torch.sum((pred_labels[:,i]==0) & (labels[:,i]==1))\n    confusion_matrix[i,1,0] += fn\n    \n    tn=torch.sum((pred_labels[:,i]==0) & (labels[:,i]==0))\n    confusion_matrix[i,1,1] += tn\n    \n  return confusion_matrix","metadata":{"id":"1AgVWoxici60"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def precisionRecallF1ConfusionMatrix(confusion_matrix):\n    for i in range(19):\n        tp = confusion_matrix[i,0,0]\n        fp = confusion_matrix[i,0,1]\n        fn = confusion_matrix[i,1,0]\n        tn = confusion_matrix[i,1,1]\n        precision = tp/(tp+fp)\n        recall = tp/(tp+fn)\n        \n        if torch.isnan(precision):\n            precision = 0\n        if torch.isnan(recall):\n            recall = 0\n            \n        epsilon=0.00001\n        f1 = 2*precision*recall/(precision+recall+ epsilon)\n\n        if torch.isnan(f1):\n            f1 = 0\n            \n        print(\"Class\",Classes[i])\n        print(\"Label number: \",Label_number[i])\n        print(\"Precision: \", precision)\n        print(\"Recall: \", recall)\n        print(\"F1: \", )\n        cf = confusion_matrix[i].detach().cpu().numpy()\n        df_cm = pd.DataFrame(cf,index=['P','N'],columns=['P','N'])\n        plt.figure(figsize = (10,7))\n        sn.heatmap(df_cm, annot=True,cmap=\"OrRd\")\n        plt.show()","metadata":{"id":"Vnmi7tViE69t"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import roc_curve, auc\ndef show_roc_curve(labels,scores):\n    plt.figure()\n    for i in range(19):\n        cur_labels = labels[:,i]\n        cur_scores = scores[:,i]\n        \n        #Calculate roc_auc score\n        fpr, tpr, thresholds = roc_curve(cur_labels, cur_scores, pos_label=1)\n        roc_auc = auc(fpr,tpr) \n\n        #+3 to ignore first colour set for background channel of masks\n        r = Colours[i*3+0+3]/255\n        g = Colours[i*3+1+3]/255\n        b = Colours[i*3+2+3]/255\n        lw=2\n        plt.plot(fpr, tpr, color=(r, g, b),lw=lw, label='ROC-AUC: '+Classes[i] + '(area = %0.2f)' % roc_auc)\n        \n        if i==4 or i==9 or i==14 or i==18:\n            plt.plot([0, 1], [0, 1], color='navy', lw=lw, linestyle='--')\n            plt.xlim([0.0, 1.0])\n            plt.ylim([0.0, 1.05])\n            plt.xlabel('False Positive Rate')\n            plt.ylabel('True Positive Rate')\n            plt.title(Classes[i]+':'+Label_number[i])\n            plt.legend(loc=\"lower right\")\n            plt.show()\n            plt.figure()","metadata":{"id":"N93-KsREc1_p"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n  test_loss, confusion_matrix, labels_arr, scores_arr = evaluation(net,test_loader)","metadata":{"id":"b9N27xihcrgG"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"precisionRecallF1ConfusionMatrix(confusion_matrix)","metadata":{"id":"qFaI6WxOcrlO"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_roc_curve(labels_arr,scores_arr)","metadata":{"id":"OZTt09o8croJ"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"FzBBZ_-gEu9M"},"execution_count":null,"outputs":[]}]}