{"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 numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport torch\nfrom torch.utils.data import Dataset,DataLoader\nimport os\nimport cv2\nimport glob\nfrom torch.autograd import Variable","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:25.172493Z","iopub.execute_input":"2022-08-15T06:33:25.173076Z","iopub.status.idle":"2022-08-15T06:33:25.178751Z","shell.execute_reply.started":"2022-08-15T06:33:25.173042Z","shell.execute_reply":"2022-08-15T06:33:25.177449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DATASET","metadata":{}},{"cell_type":"code","source":"class Hubmapdataset(Dataset):\n    def __init__(self,mode):\n        #Get image from directory and sort them \n        # join image folder and  image then save in a list \n        #for ex :['../input/hubmaphacking2022-comp-256x256/train/10044_0000.png','__/10044_0001.png'...]\n        self.image_dir = sorted(glob.glob('../input/hubmaphacking2022-comp-256x256/train/*'))\n        # join mask folder and mask then save in a list \n        #for ex :['../input/hubmaphacking2022-comp-256x256/mask/10044_0000.png','__/10044_0001.png'...]\n        self.mask_dir=sorted(glob.glob('../input/hubmaphacking2022-comp-256x256/masks/*'))\n        self.mode=mode\n        assert mode in ['train','test']\n        #Creat mode train and test to do augmentation\n        if self.mode=='train':\n            self.augmentation=get_training_augmentation()                  \n            \n    def __getitem__(self, i):\n        # read image file\n        image = cv2.imread(self.image_dir[i])\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        # read mask file\n        mask = cv2.imread(self.mask_dir[i],0)\n        mask=np.expand_dims(mask, axis=2)\n        # apply augmentations\n        if self.mode=='train':\n            #do augmentation for image and mask\n            sample = self.augmentation(image=image, mask=mask)\n            image = sample['image'].astype(np.float32)/255\n            image=image.transpose(2, 0, 1)\n            mask=sample['mask'].astype(np.float32).transpose(2,1,0)\n            return torch.tensor(image), torch.tensor(mask)\n        elif self.mode=='test':\n            image = image.astype(np.float32)/255\n            image=image.transpose(2, 0, 1)\n            mask=mask.transpose(2,1,0)\n            return torch.tensor(image),torch.tensor(mask)\n    \n    def __len__(self):\n        return len(self.image_dir)","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:25.220378Z","iopub.execute_input":"2022-08-15T06:33:25.221034Z","iopub.status.idle":"2022-08-15T06:33:25.232539Z","shell.execute_reply.started":"2022-08-15T06:33:25.221005Z","shell.execute_reply":"2022-08-15T06:33:25.231491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DATA AUGMENTATION","metadata":{}},{"cell_type":"code","source":"import albumentations as albu\ndef get_training_augmentation(): \n    train_transform = [\n        albu.HorizontalFlip(p=0.5),\n        #albu.ShiftScaleRotate(scale_limit=0.5, rotate_limit=0, shift_limit=0.1, p=1, border_mode=0),\n        albu.OneOf(\n            [\n                albu.CLAHE(p=1),\n                albu.RandomBrightness(p=1),\n                albu.RandomGamma(p=1),\n            ],\n            p=0.9,),\n        albu.OneOf(\n            [\n                albu.RandomContrast(p=0.5),\n                albu.HueSaturationValue(p=0.5),\n            ],\n            p=0.9,),\n         albu.OneOf([\n            albu.OpticalDistortion(p=0.3),\n            albu.GridDistortion(p=.1),\n            albu.IAAPiecewiseAffine(p=0.3),\n        ], p=0.3),]\n    return albu.Compose(train_transform)","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:25.248922Z","iopub.execute_input":"2022-08-15T06:33:25.249199Z","iopub.status.idle":"2022-08-15T06:33:26.613228Z","shell.execute_reply.started":"2022-08-15T06:33:25.249174Z","shell.execute_reply":"2022-08-15T06:33:26.612288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_dir='../input/hubmaphacking2022-comp-256x256/train'\nmask_dir='../input/hubmaphacking2022-comp-256x256/mask'\ntrain_dataset=Hubmapdataset(mode='train')\ntest_dataset=Hubmapdataset(mode='test')\ntrain_loader=DataLoader(dataset=train_dataset,\n                                  batch_size=16,\n                                  shuffle=True,\n                                  num_workers=4,\n                                  pin_memory=True)\ntest_loader=DataLoader(dataset=test_dataset,\n                                  batch_size=16,\n                                  shuffle=False,\n                                  num_workers=4,\n                                  pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:26.615196Z","iopub.execute_input":"2022-08-15T06:33:26.615548Z","iopub.status.idle":"2022-08-15T06:33:26.879709Z","shell.execute_reply.started":"2022-08-15T06:33:26.615513Z","shell.execute_reply":"2022-08-15T06:33:26.878512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# HELPER LOSS FUNCTION","metadata":{}},{"cell_type":"code","source":"def symmetric_lovasz(outputs, targets):\n    return 0.5*(lovasz_hinge(outputs, targets) + lovasz_hinge(-outputs, 1.0 - targets))\ndef lovasz_grad(gt_sorted):\n    \"\"\"\n    Computes gradient of the Lovasz extension w.r.t sorted errors\n    See Alg. 1 in paper\n    \"\"\"\n    p = len(gt_sorted)\n    gts = gt_sorted.sum()\n    intersection = gts - gt_sorted.float().cumsum(0)\n    union = gts + (1 - gt_sorted).float().cumsum(0)\n    jaccard = 1. - intersection / union\n    if p > 1: # cover 1-pixel case\n        jaccard[1:p] = jaccard[1:p] - jaccard[0:-1]\n    return jaccard\ndef mean(l, ignore_nan=False, empty=0):\n    \"\"\"\n    nanmean compatible with generators.\n    \"\"\"\n    l = iter(l)\n    if ignore_nan:\n        l = ifilterfalse(np.isnan, l)\n    try:\n        n = 1\n        acc = next(l)\n    except StopIteration:\n        if empty == 'raise':\n            raise ValueError('Empty mean')\n        return empty\n    for n, v in enumerate(l, 2):\n        acc += v\n    if n == 1:\n        return acc\n    return acc / n\ndef lovasz_hinge_flat(logits, labels):\n    \"\"\"\n    Binary Lovasz hinge loss\n      logits: [P] Variable, logits at each prediction (between -\\infty and +\\infty)\n      labels: [P] Tensor, binary ground truth labels (0 or 1)\n      ignore: label to ignore\n    \"\"\"\n    if len(labels) == 0:\n        # only void pixels, the gradients should be 0\n        return logits.sum() * 0.\n    signs = 2. * labels.float() - 1.\n    errors = (1. - logits * Variable(signs))\n    errors_sorted, perm = torch.sort(errors, dim=0, descending=True)\n    perm = perm.data\n    gt_sorted = labels[perm]\n    grad = lovasz_grad(gt_sorted)\n    loss = torch.dot(F.relu(errors_sorted), Variable(grad))\n    return loss\ndef flatten_binary_scores(scores, labels, ignore=None):\n    \"\"\"\n    Flattens predictions in the batch (binary case)\n    Remove labels equal to 'ignore'\n    \"\"\"\n    scores = scores.view(-1)\n    labels = labels.view(-1)\n    if ignore is None:\n        return scores, labels\n    valid = (labels != ignore)\n    vscores = scores[valid]\n    vlabels = labels[valid]\n    return vscores, vlabels\ndef lovasz_hinge(logits, labels, per_image=True, ignore=None):\n    \"\"\"\n    Binary Lovasz hinge loss\n      logits: [B, H, W] Variable, logits at each pixel (between -\\infty and +\\infty)\n      labels: [B, H, W] Tensor, binary ground truth masks (0 or 1)\n      per_image: compute the loss per image instead of per batch\n      ignore: void class id\n    \"\"\"\n    if per_image:\n        loss = mean(lovasz_hinge_flat(*flatten_binary_scores(log.unsqueeze(0), lab.unsqueeze(0), ignore))\n                          for log, lab in zip(logits, labels))\n    else:\n        loss = lovasz_hinge_flat(*flatten_binary_scores(logits, labels, ignore))\n    return loss","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:26.881269Z","iopub.execute_input":"2022-08-15T06:33:26.881891Z","iopub.status.idle":"2022-08-15T06:33:26.896471Z","shell.execute_reply.started":"2022-08-15T06:33:26.881854Z","shell.execute_reply":"2022-08-15T06:33:26.895599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nclass BasicConv2d(nn.Module):\n    def __init__(self, in_planes, out_planes, kernel_size, stride=1, padding=0, dilation=1):\n        super(BasicConv2d, self).__init__()\n        self.conv = nn.Conv2d(in_planes, out_planes,\n                              kernel_size=kernel_size, stride=stride,\n                              padding=padding, dilation=dilation, bias=False)\n        self.bn = nn.BatchNorm2d(out_planes)\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.bn(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:26.900992Z","iopub.execute_input":"2022-08-15T06:33:26.901257Z","iopub.status.idle":"2022-08-15T06:33:26.913582Z","shell.execute_reply.started":"2022-08-15T06:33:26.901232Z","shell.execute_reply":"2022-08-15T06:33:26.912604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# RECEPTIVE FIELD BLOCK\n","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nclass RFB_modified(nn.Module):\n    def __init__(self, in_channel, out_channel):\n        super(RFB_modified, self).__init__()\n        self.relu = nn.ReLU(True)\n        self.branch0 = nn.Sequential(\n            BasicConv2d(in_channel, out_channel, 1),\n        )\n        self.branch1 = nn.Sequential(\n            BasicConv2d(in_channel, out_channel, 1),\n            BasicConv2d(out_channel, out_channel, kernel_size=(1, 3), padding=(0, 1)),\n            BasicConv2d(out_channel, out_channel, kernel_size=(3, 1), padding=(1, 0)),\n            BasicConv2d(out_channel, out_channel, 3, padding=3, dilation=3)\n        )\n        self.branch2 = nn.Sequential(\n            BasicConv2d(in_channel, out_channel, 1),\n            BasicConv2d(out_channel, out_channel, kernel_size=(1, 5), padding=(0, 2)),\n            BasicConv2d(out_channel, out_channel, kernel_size=(5, 1), padding=(2, 0)),\n            BasicConv2d(out_channel, out_channel, 3, padding=5, dilation=5)\n        )\n        self.branch3 = nn.Sequential(\n            BasicConv2d(in_channel, out_channel, 1),\n            BasicConv2d(out_channel, out_channel, kernel_size=(1, 7), padding=(0, 3)),\n            BasicConv2d(out_channel, out_channel, kernel_size=(7, 1), padding=(3, 0)),\n            BasicConv2d(out_channel, out_channel, 3, padding=7, dilation=7)\n        )\n        self.conv_cat = BasicConv2d(4*out_channel, out_channel, 3, padding=1)\n        self.conv_res = BasicConv2d(in_channel, out_channel, 1)\n\n    def forward(self, x):\n        x0 = self.branch0(x)\n        x1 = self.branch1(x)\n        x2 = self.branch2(x)\n        x3 = self.branch3(x)\n        x_cat = self.conv_cat(torch.cat((x0, x1, x2, x3), 1))\n\n        x = self.relu(x_cat + self.conv_res(x))\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:26.915223Z","iopub.execute_input":"2022-08-15T06:33:26.915992Z","iopub.status.idle":"2022-08-15T06:33:26.930661Z","shell.execute_reply.started":"2022-08-15T06:33:26.915956Z","shell.execute_reply":"2022-08-15T06:33:26.929609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# PARTIAL DECODER","metadata":{}},{"cell_type":"code","source":"class aggregation(nn.Module):\n    # dense aggregation, it can be replaced by other aggregation previous, such as DSS, amulet, and so on.\n    # used after MSF\n    def __init__(self, channel):\n        super(aggregation, self).__init__()\n        self.relu = nn.ReLU(True)\n\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv_upsample1 = BasicConv2d(channel, channel, 3, padding=1)\n        self.conv_upsample2 = BasicConv2d(channel, channel, 3, padding=1)\n        self.conv_upsample3 = BasicConv2d(channel, channel, 3, padding=1)\n        self.conv_upsample4 = BasicConv2d(channel, channel, 3, padding=1)\n        self.conv_upsample5 = BasicConv2d(2*channel, 2*channel, 3, padding=1)\n\n        self.conv_concat2 = BasicConv2d(2*channel, 2*channel, 3, padding=1)\n        self.conv_concat3 = BasicConv2d(3*channel, 3*channel, 3, padding=1)\n        self.conv4 = BasicConv2d(3*channel, 3*channel, 3, padding=1)\n        self.conv5 = nn.Conv2d(3*channel, 1, 1)\n\n    def forward(self, x1, x2, x3):\n        x1_1 = x1\n        x2_1 = self.conv_upsample1(self.upsample(x1)) * x2\n        x3_1 = self.conv_upsample2(self.upsample(self.upsample(x1))) \\\n               * self.conv_upsample3(self.upsample(x2)) * x3\n\n        x2_2 = torch.cat((x2_1, self.conv_upsample4(self.upsample(x1_1))), 1)\n        x2_2 = self.conv_concat2(x2_2)\n\n        x3_2 = torch.cat((x3_1, self.conv_upsample5(self.upsample(x2_2))), 1)\n        x3_2 = self.conv_concat3(x3_2)\n\n        x = self.conv4(x3_2)\n        x = self.conv5(x)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:26.932324Z","iopub.execute_input":"2022-08-15T06:33:26.932844Z","iopub.status.idle":"2022-08-15T06:33:26.947178Z","shell.execute_reply.started":"2022-08-15T06:33:26.932810Z","shell.execute_reply":"2022-08-15T06:33:26.946273Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# RESNET BACKBONE","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport math\nimport torch.utils.model_zoo as model_zoo\nimport torch\nimport torch.nn.functional as F\n\n__all__ = ['Res2Net', 'res2net50_v1b', 'res2net101_v1b', 'res2net50_v1b_26w_4s']\n\nmodel_urls = {\n    'res2net50_v1b_26w_4s': 'https://shanghuagao.oss-cn-beijing.aliyuncs.com/res2net/res2net50_v1b_26w_4s-3cf99910.pth',\n    'res2net101_v1b_26w_4s': 'https://shanghuagao.oss-cn-beijing.aliyuncs.com/res2net/res2net101_v1b_26w_4s-0812c246.pth',\n}\nclass Bottle2neck(nn.Module):\n    expansion = 4\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None, baseWidth=26, scale=4, stype='normal'):\n        \"\"\" Constructor\n        Args:\n            inplanes: input channel dimensionality\n            planes: output channel dimensionality\n            stride: conv stride. Replaces pooling layer.\n            downsample: None when stride = 1\n            baseWidth: basic width of conv3x3\n            scale: number of scale.\n            type: 'normal': normal set. 'stage': first block of a new stage.\n        \"\"\"\n        super(Bottle2neck, self).__init__()\n\n        width = int(math.floor(planes * (baseWidth / 64.0)))\n        self.conv1 = nn.Conv2d(inplanes, width * scale, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(width * scale)\n\n        if scale == 1:\n            self.nums = 1\n        else:\n            self.nums = scale - 1\n        if stype == 'stage':\n            self.pool = nn.AvgPool2d(kernel_size=3, stride=stride, padding=1)\n        convs = []\n        bns = []\n        for i in range(self.nums):\n            convs.append(nn.Conv2d(width, width, kernel_size=3, stride=stride, padding=1, bias=False))\n            bns.append(nn.BatchNorm2d(width))\n        self.convs = nn.ModuleList(convs)\n        self.bns = nn.ModuleList(bns)\n\n        self.conv3 = nn.Conv2d(width * scale, planes * self.expansion, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm2d(planes * self.expansion)\n\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n        self.stype = stype\n        self.scale = scale\n        self.width = width\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        spx = torch.split(out, self.width, 1)\n        for i in range(self.nums):\n            if i == 0 or self.stype == 'stage':\n                sp = spx[i]\n            else:\n                sp = sp + spx[i]\n            sp = self.convs[i](sp)\n            sp = self.relu(self.bns[i](sp))\n            if i == 0:\n                out = sp\n            else:\n                out = torch.cat((out, sp), 1)\n        if self.scale != 1 and self.stype == 'normal':\n            out = torch.cat((out, spx[self.nums]), 1)\n        elif self.scale != 1 and self.stype == 'stage':\n            out = torch.cat((out, self.pool(spx[self.nums])), 1)\n\n        out = self.conv3(out)\n        out = self.bn3(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\nclass Res2Net(nn.Module):\n\n    def __init__(self, block, layers, baseWidth=26, scale=4, num_classes=1000):\n        self.inplanes = 64\n        super(Res2Net, self).__init__()\n        self.baseWidth = baseWidth\n        self.scale = scale\n        self.conv1 = nn.Sequential(\n            nn.Conv2d(3, 32, 3, 2, 1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, 32, 3, 1, 1, bias=False),\n            nn.BatchNorm2d(32),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(32, 64, 3, 1, 1, bias=False)\n        )\n        self.bn1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU()\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        self.layer1 = self._make_layer(block, 64, layers[0])\n        self.layer2 = self._make_layer(block, 128, layers[1], stride=2)\n        self.layer3 = self._make_layer(block, 256, layers[2], stride=2)\n        self.layer4 = self._make_layer(block, 512, layers[3], stride=2)\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(512 * block.expansion, num_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n                \n    def _make_layer(self, block, planes, blocks, stride=1):\n        downsample = None\n        if stride != 1 or self.inplanes != planes * block.expansion:\n            downsample = nn.Sequential(\n                nn.AvgPool2d(kernel_size=stride, stride=stride,\n                             ceil_mode=True, count_include_pad=False),\n                nn.Conv2d(self.inplanes, planes * block.expansion,\n                          kernel_size=1, stride=1, bias=False),\n                nn.BatchNorm2d(planes * block.expansion),\n            )\n\n        layers = []\n        layers.append(block(self.inplanes, planes, stride, downsample=downsample,\n                            stype='stage', baseWidth=self.baseWidth, scale=self.scale))\n        self.inplanes = planes * block.expansion\n        for i in range(1, blocks):\n            layers.append(block(self.inplanes, planes, baseWidth=self.baseWidth, scale=self.scale))\n\n        return nn.Sequential(*layers)\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        x = self.avgpool(x)\n        x = x.view(x.size(0), -1)\n        x = self.fc(x)\n\n        return x\ndef res2net50_v1b_26w_4s(pretrained=True, **kwargs):\n    \"\"\"Constructs a Res2Net-50_v1b_26w_4s lib.\n    Args:\n        pretrained (bool): If True, returns a lib pre-trained on ImageNet\n    \"\"\"\n    model = Res2Net(Bottle2neck, [3, 4, 6, 3], baseWidth=26, scale=4, **kwargs)\n    if pretrained:\n        model_state = torch.load('../input/res2net50/res2net50_v1b_26w_4s-3cf99910.pth')\n        model.load_state_dict(model_state)\n        # lib.load_state_dict(model_zoo.load_url(model_urls['res2net50_v1b_26w_4s']))\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:26.949138Z","iopub.execute_input":"2022-08-15T06:33:26.949873Z","iopub.status.idle":"2022-08-15T06:33:26.981435Z","shell.execute_reply.started":"2022-08-15T06:33:26.949838Z","shell.execute_reply":"2022-08-15T06:33:26.980385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PraNet(nn.Module):\n    # res2net based encoder decoder\n    def __init__(self, channel=32):\n        super(PraNet, self).__init__()\n        # ---- ResNet Backbone ----\n        self.resnet = res2net50_v1b_26w_4s(pretrained=True)\n        # ---- Receptive Field Block like module ----\n        self.rfb2_1 = RFB_modified(512, channel)\n        self.rfb3_1 = RFB_modified(1024, channel)\n        self.rfb4_1 = RFB_modified(2048, channel)\n        # ---- Partial Decoder ----\n        self.agg1 = aggregation(channel)\n        # ---- reverse attention branch 4 ----\n        self.ra4_conv1 = BasicConv2d(2048, 256, kernel_size=1)\n        self.ra4_conv2 = BasicConv2d(256, 256, kernel_size=5, padding=2)\n        self.ra4_conv3 = BasicConv2d(256, 256, kernel_size=5, padding=2)\n        self.ra4_conv4 = BasicConv2d(256, 256, kernel_size=5, padding=2)\n        self.ra4_conv5 = BasicConv2d(256, 1, kernel_size=1)\n        # ---- reverse attention branch 3 ----\n        self.ra3_conv1 = BasicConv2d(1024, 64, kernel_size=1)\n        self.ra3_conv2 = BasicConv2d(64, 64, kernel_size=3, padding=1)\n        self.ra3_conv3 = BasicConv2d(64, 64, kernel_size=3, padding=1)\n        self.ra3_conv4 = BasicConv2d(64, 1, kernel_size=3, padding=1)\n        # ---- reverse attention branch 2 ----\n        self.ra2_conv1 = BasicConv2d(512, 64, kernel_size=1)\n        self.ra2_conv2 = BasicConv2d(64, 64, kernel_size=3, padding=1)\n        self.ra2_conv3 = BasicConv2d(64, 64, kernel_size=3, padding=1)\n        self.ra2_conv4 = BasicConv2d(64, 1, kernel_size=3, padding=1)\n\n    def forward(self, x):\n        x = self.resnet.conv1(x)\n        x = self.resnet.bn1(x)\n        x = self.resnet.relu(x)\n        x = self.resnet.maxpool(x)      # bs, 64, 88, 88\n        # ---- low-level features ----\n        x1 = self.resnet.layer1(x)      # bs, 256, 88, 88\n        x2 = self.resnet.layer2(x1)     # bs, 512, 44, 44\n\n        x3 = self.resnet.layer3(x2)     # bs, 1024, 22, 22\n        x4 = self.resnet.layer4(x3)     # bs, 2048, 11, 11\n        x2_rfb = self.rfb2_1(x2)        # channel -> 32\n        x3_rfb = self.rfb3_1(x3)        # channel -> 32\n        x4_rfb = self.rfb4_1(x4)        # channel -> 32\n\n        ra5_feat = self.agg1(x4_rfb, x3_rfb, x2_rfb)\n        lateral_map_5 = F.interpolate(ra5_feat, scale_factor=8, mode='bilinear')    # NOTES: Sup-1 (bs, 1, 44, 44) -> (bs, 1, 352, 352)\n\n        # ---- reverse attention branch_4 ----\n        crop_4 = F.interpolate(ra5_feat, scale_factor=0.25, mode='bilinear')\n        x = -1*(torch.sigmoid(crop_4)) + 1\n        x = x.expand(-1, 2048, -1, -1).mul(x4)\n        x = self.ra4_conv1(x)\n        x = F.relu(self.ra4_conv2(x))\n        x = F.relu(self.ra4_conv3(x))\n        x = F.relu(self.ra4_conv4(x))\n        ra4_feat = self.ra4_conv5(x)\n        x = ra4_feat + crop_4\n        lateral_map_4 = F.interpolate(x, scale_factor=32, mode='bilinear')  # NOTES: Sup-2 (bs, 1, 11, 11) -> (bs, 1, 352, 352)\n\n        # ---- reverse attention branch_3 ----\n        crop_3 = F.interpolate(x, scale_factor=2, mode='bilinear')\n        x = -1*(torch.sigmoid(crop_3)) + 1\n        x = x.expand(-1, 1024, -1, -1).mul(x3)\n        x = self.ra3_conv1(x)\n        x = F.relu(self.ra3_conv2(x))\n        x = F.relu(self.ra3_conv3(x))\n        ra3_feat = self.ra3_conv4(x)\n        x = ra3_feat + crop_3\n        lateral_map_3 = F.interpolate(x, scale_factor=16, mode='bilinear')  # NOTES: Sup-3 (bs, 1, 22, 22) -> (bs, 1, 352, 352)\n\n        # ---- reverse attention branch_2 ----\n        crop_2 = F.interpolate(x, scale_factor=2, mode='bilinear')\n        x = -1*(torch.sigmoid(crop_2)) + 1\n        x = x.expand(-1, 512, -1, -1).mul(x2)\n        x = self.ra2_conv1(x)\n        x = F.relu(self.ra2_conv2(x))\n        x = F.relu(self.ra2_conv3(x))\n        ra2_feat = self.ra2_conv4(x)\n        x = ra2_feat + crop_2\n        lateral_map_2 = F.interpolate(x, scale_factor=8, mode='bilinear')   # NOTES: Sup-4 (bs, 1, 44, 44) -> (bs, 1, 352, 352)\n\n        return lateral_map_5, lateral_map_4, lateral_map_3, lateral_map_2","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:26.985355Z","iopub.execute_input":"2022-08-15T06:33:26.986512Z","iopub.status.idle":"2022-08-15T06:33:27.027457Z","shell.execute_reply.started":"2022-08-15T06:33:26.986462Z","shell.execute_reply":"2022-08-15T06:33:27.025328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def do_kaggle_metric(predict,truth, threshold=0.5):\n\n    N = len(predict)\n    predict = predict.reshape(N,-1)\n    truth   = truth.reshape(N,-1)\n\n    predict = predict>threshold\n    truth   = truth>0.5\n    intersection = truth & predict\n    union        = truth | predict\n    iou = intersection.sum(1)/(union.sum(1)+1e-8)\n\n    #-------------------------------------------\n    result = []\n    precision = []\n    is_empty_truth   = (truth.sum(1)==0)\n    is_empty_predict = (predict.sum(1)==0)\n\n    threshold = np.array([0.50, 0.55, 0.60, 0.65, 0.70, 0.75, 0.80, 0.85, 0.90, 0.95])\n    for t in threshold:\n        p = iou>=t\n\n        tp  = (~is_empty_truth)  & (~is_empty_predict) & (iou> t)\n        fp  = (~is_empty_truth)  & (~is_empty_predict) & (iou<=t)\n        fn  = (~is_empty_truth)  & ( is_empty_predict)\n        fp_empty = ( is_empty_truth)  & (~is_empty_predict)\n        tn_empty = ( is_empty_truth)  & ( is_empty_predict)\n\n        p = (tp + tn_empty) / (tp + tn_empty + fp + fp_empty + fn)\n\n        result.append( np.column_stack((tp,fp,fn,tn_empty,fp_empty)) )\n        precision.append(p)\n\n    result = np.array(result).transpose(1,2,0)\n    precision = np.column_stack(precision)\n    precision = precision.mean(1)\n\n    return precision, result, threshold","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:27.033588Z","iopub.execute_input":"2022-08-15T06:33:27.034330Z","iopub.status.idle":"2022-08-15T06:33:27.054885Z","shell.execute_reply.started":"2022-08-15T06:33:27.034284Z","shell.execute_reply":"2022-08-15T06:33:27.053795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_pranet(test_loader, model):\n    running_loss = 0.0\n    predicts = []\n    truths = []\n    model.eval()\n    for inputs, masks in test_loader:\n        inputs, masks = inputs.to(device), masks.to(device)\n        with torch.set_grad_enabled(False):\n            lateral_map_5, lateral_map_4, lateral_map_3, lateral_map_2= model(inputs)\n            #loss = lovasz_hinge(outputs.squeeze(1), masks.squeeze(1))\n            loss5=symmetric_lovasz(lateral_map_5.squeeze(1), masks.squeeze(1))\n            loss4=symmetric_lovasz(lateral_map_4.squeeze(1), masks.squeeze(1))\n            loss3=symmetric_lovasz(lateral_map_3.squeeze(1), masks.squeeze(1))\n            loss2=symmetric_lovasz(lateral_map_2.squeeze(1), masks.squeeze(1))\n            loss=loss5+loss4+loss3+loss2\n            outputs=lateral_map_2#+lateral_map_4+lateral_map_3+lateral_map_5\n        predicts.append(F.sigmoid(outputs).detach().cpu().numpy())\n        truths.append(masks.detach().cpu().numpy())\n        running_loss += loss.item() * inputs.size(0)\n\n    predicts = np.concatenate(predicts).squeeze()\n    truths = np.concatenate(truths).squeeze()\n    precision, _, _ = do_kaggle_metric(predicts, truths, 0.5)\n    precision = precision.mean()\n    epoch_loss = running_loss / test_dataset.__len__()\n    return epoch_loss, precision,predicts\n\n\ndef train_pranet(train_loader, model):\n    running_loss = 0.0\n    data_size = train_dataset.__len__()\n    model.train()\n    # for inputs, masks, labels in progress_bar(train_loader, parent=mb):\n    for inputs, masks in train_loader:\n        inputs, masks = inputs.to(device), masks.to(device)\n        optimizer.zero_grad()\n        with torch.set_grad_enabled(True):\n            lateral_map_5, lateral_map_4, lateral_map_3, lateral_map_2= model(inputs)\n            #loss = lovasz_hinge(outputs.squeeze(1), masks.squeeze(1))\n            loss5=symmetric_lovasz(lateral_map_5.squeeze(1), masks.squeeze(1))\n            loss4=symmetric_lovasz(lateral_map_4.squeeze(1), masks.squeeze(1))\n            loss3=symmetric_lovasz(lateral_map_3.squeeze(1), masks.squeeze(1))\n            loss2=symmetric_lovasz(lateral_map_2.squeeze(1), masks.squeeze(1))\n            loss=loss5+loss4+loss3+loss2\n            #https://discuss.pytorch.org/t/what-does-the-backward-function-do/9944/2\n            loss.backward()\n            optimizer.step()\n\n        running_loss += loss.item() * inputs.size(0)\n        # mb.child.comment = 'loss: {}'.format(loss.item())\n    epoch_loss = running_loss / data_size\n    return epoch_loss ","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:27.060890Z","iopub.execute_input":"2022-08-15T06:33:27.061241Z","iopub.status.idle":"2022-08-15T06:33:27.086917Z","shell.execute_reply.started":"2022-08-15T06:33:27.061209Z","shell.execute_reply":"2022-08-15T06:33:27.081081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model=PraNet()\ndevice = torch.device('cuda' if True else 'cpu')\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:27.088382Z","iopub.execute_input":"2022-08-15T06:33:27.088817Z","iopub.status.idle":"2022-08-15T06:33:31.961629Z","shell.execute_reply.started":"2022-08-15T06:33:27.088778Z","shell.execute_reply":"2022-08-15T06:33:31.960711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_acc=0\nnum_snapshot=0\nscheduler_step = 20 // 5\n#learning rate tốt: max_lr=1e-04\nmax_lr=1e-04\nmin_lr=1e-07\noptimizer = torch.optim.Adam(model.parameters(),max_lr)#,weight_decay=1e-4)  #\n# optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9,\n                                     #weight_decay=1e-4)\nlr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, scheduler_step,min_lr)","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:31.963065Z","iopub.execute_input":"2022-08-15T06:33:31.963702Z","iopub.status.idle":"2022-08-15T06:33:31.972251Z","shell.execute_reply.started":"2022-08-15T06:33:31.963665Z","shell.execute_reply":"2022-08-15T06:33:31.971317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(150):\n    train_loss = train_pranet(train_loader, model)\n    val_loss, accuracy,pred = test_pranet(test_loader, model)\n    lr_scheduler.step()\n\n    if accuracy > best_acc:\n        best_acc = accuracy\n        best_param = model.state_dict()\n\n    if (epoch + 1) % scheduler_step == 0:\n#         torch.save(best_param, path)\n#         optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9,weight_decay=1e-4)\n        optimizer = torch.optim.Adam(model.parameters(),max_lr)#,weight_decay=1e-4)\n        lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, scheduler_step, min_lr)\n        num_snapshot += 1\n        best_acc = 0\n\n    # mb.write('epoch: {} train_loss: {:.3f} val_loss: {:.3f} val_accuracy: {:.3f}'.format(epoch + 1, train_loss,\n    #                                                                                     val_loss, accuracy))\n    print('epoch: {} train_loss: {:.3f} val_loss: {:.3f} val_accuracy: {:.3f}'.format(epoch + 1, train_loss,\n                                                                                      val_loss, accuracy))","metadata":{"execution":{"iopub.status.busy":"2022-08-15T06:33:31.973757Z","iopub.execute_input":"2022-08-15T06:33:31.974154Z","iopub.status.idle":"2022-08-15T11:36:29.445683Z","shell.execute_reply.started":"2022-08-15T06:33:31.974119Z","shell.execute_reply":"2022-08-15T11:36:29.442147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(50):\n    train_loss = train_pranet(train_loader, model)\n    val_loss, accuracy,pred = test_pranet(test_loader, model)\n    lr_scheduler.step()\n\n    if accuracy > best_acc:\n        best_acc = accuracy\n        best_param = model.state_dict()\n\n    if (epoch + 1) % scheduler_step == 0:\n#         torch.save(best_param, path)\n#         optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9,weight_decay=1e-4)\n        optimizer = torch.optim.Adam(model.parameters(),max_lr)#,weight_decay=1e-4)\n        lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, scheduler_step, min_lr)\n        num_snapshot += 1\n        best_acc = 0\n\n    # mb.write('epoch: {} train_loss: {:.3f} val_loss: {:.3f} val_accuracy: {:.3f}'.format(epoch + 1, train_loss,\n    #                                                                                     val_loss, accuracy))\n    print('epoch: {} train_loss: {:.3f} val_loss: {:.3f} val_accuracy: {:.3f}'.format(epoch + 1, train_loss,\n                                                                                      val_loss, accuracy))","metadata":{"execution":{"iopub.status.busy":"2022-08-15T11:37:42.479490Z","iopub.execute_input":"2022-08-15T11:37:42.479992Z","iopub.status.idle":"2022-08-15T13:43:42.862832Z","shell.execute_reply.started":"2022-08-15T11:37:42.479952Z","shell.execute_reply":"2022-08-15T13:43:42.861787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save('pred', pred)","metadata":{"execution":{"iopub.status.busy":"2022-08-15T14:51:16.131514Z","iopub.execute_input":"2022-08-15T14:51:16.132277Z","iopub.status.idle":"2022-08-15T14:51:16.694417Z","shell.execute_reply.started":"2022-08-15T14:51:16.132230Z","shell.execute_reply":"2022-08-15T14:51:16.693386Z"},"trusted":true},"execution_count":null,"outputs":[]}]}