{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load in \n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the \"../input/\" directory.\n# For example, running this (by clicking run or pressing Shift+Enter) will list the files in the input directory\n\nimport os\nprint(os.listdir(\"../input\"))\n\n# Any results you write to the current directory are saved as output.","execution_count":1,"outputs":[{"output_type":"stream","text":"['pytorch-model-zoo', 'pretrainednasnetpytorch', 'imet-2019-fgvc6']\n","name":"stdout"}]},{"metadata":{},"cell_type":"markdown","source":"## Layers of NASNet"},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":true},"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\n\nclass SeparableConv2d(nn.Sequential):\n\n    def __init__(self, in_channels, out_channels, dw_kernel, dw_stride, dw_padding, bias=False):\n        super(SeparableConv2d, self).__init__()\n        self.depthwise_conv2d = nn.Conv2d(in_channels, in_channels, dw_kernel,\n                                          stride=dw_stride,\n                                          padding=dw_padding,\n                                          bias=bias,\n                                          groups=in_channels)\n        self.pointwise_conv2d = nn.Conv2d(in_channels, out_channels, 1, stride=1, bias=bias)\n\n\nclass BranchSeparables(nn.Sequential):\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, bias=False):\n        super(BranchSeparables, self).__init__()\n        self.relu = nn.ReLU()\n        self.separable_1 = SeparableConv2d(in_channels, in_channels, kernel_size, stride, padding, bias=bias)\n        self.bn_sep_1 = nn.BatchNorm2d(in_channels, eps=0.001, momentum=0.1, affine=True)\n        self.relu1 = nn.ReLU()\n        self.separable_2 = SeparableConv2d(in_channels, out_channels, kernel_size, 1, padding, bias=bias)\n        self.bn_sep_2 = nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True)\n\n\nclass BranchSeparablesStem(nn.Sequential):\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, bias=False):\n        super(BranchSeparablesStem, self).__init__()\n        self.relu = nn.ReLU()\n        self.separable_1 = SeparableConv2d(in_channels, out_channels, kernel_size, stride, padding, bias=bias)\n        self.bn_sep_1 = nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True)\n        self.relu1 = nn.ReLU()\n        self.separable_2 = SeparableConv2d(out_channels, out_channels, kernel_size, 1, padding, bias=bias)\n        self.bn_sep_2 = nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True)\n\n\ndef ReductionCellBranchCombine(cell, x_left, x_right):\n\n    x_comb_iter_0_left = cell.comb_iter_0_left(x_left)\n    x_comb_iter_0_right = cell.comb_iter_0_right(x_right)\n    x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n    x_comb_iter_1_left = cell.comb_iter_1_left(x_left)\n    x_comb_iter_1_right = cell.comb_iter_1_right(x_right)\n    x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n    x_comb_iter_2_left = cell.comb_iter_2_left(x_left)\n    x_comb_iter_2_right = cell.comb_iter_2_right(x_right)\n    x_comb_iter_2 = x_comb_iter_2_left + x_comb_iter_2_right\n\n    x_comb_iter_3_right = cell.comb_iter_3_right(x_comb_iter_0)\n    x_comb_iter_3 = x_comb_iter_3_right + x_comb_iter_1\n\n    x_comb_iter_4_left = cell.comb_iter_4_left(x_comb_iter_0)\n    x_comb_iter_4_right = cell.comb_iter_4_right(x_left)\n    x_comb_iter_4 = x_comb_iter_4_left + x_comb_iter_4_right\n\n    x_out = torch.cat([x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n    return x_out\n\n\nclass CellStem0(nn.Module):\n\n    def __init__(self, in_channels, out_channels):\n        super(CellStem0, self).__init__()\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels, out_channels, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True))\n\n        self.comb_iter_0_left = BranchSeparables(out_channels, out_channels, 5, 2, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparablesStem(in_channels, out_channels, 7, 2, 3, bias=False)\n\n        self.comb_iter_1_left = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_1_right = BranchSeparablesStem(in_channels, out_channels, 7, 2, 3, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=2, padding=1, count_include_pad=False)\n        self.comb_iter_2_right = BranchSeparablesStem(in_channels, out_channels, 5, 2, 2, bias=False)\n\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels, out_channels, 3, 1, 1, bias=False)\n        self.comb_iter_4_right = nn.MaxPool2d(3, stride=2, padding=1)\n\n    def forward(self, x):\n        x1 = self.conv_1x1(x)\n\n        return ReductionCellBranchCombine(self, x1, x)\n\n\nclass CellStem1(nn.Module):\n\n    def __init__(self, in_channels_x, in_channels_h, out_channels):\n        super(CellStem1, self).__init__()\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_x, out_channels, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True))\n\n        self.relu = nn.ReLU()\n        self.path_1 = nn.Sequential()\n        self.path_1.add_module('avgpool', nn.AvgPool2d(1, stride=2, count_include_pad=False))\n        self.path_1.add_module('conv', nn.Conv2d(in_channels_h, out_channels//2, 1, stride=1, bias=False))\n        self.path_2 = nn.Sequential()\n        self.path_2.add_module('avgpool', nn.AvgPool2d(1, stride=2, ceil_mode=True, count_include_pad=False)) # ceil mode for padding\n        self.path_2.add_module('conv', nn.Conv2d(in_channels_h, out_channels//2, 1, stride=1, bias=False))\n\n        self.final_path_bn = nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True)\n\n        self.comb_iter_0_left = BranchSeparables(out_channels, out_channels, 5, 2, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparables(out_channels, out_channels, 7, 2, 3, bias=False)\n\n        self.comb_iter_1_left = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_1_right = BranchSeparables(out_channels, out_channels, 7, 2, 3, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=2, padding=1, count_include_pad=False)\n        self.comb_iter_2_right = BranchSeparables(out_channels, out_channels, 5, 2, 2, bias=False)\n\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels, out_channels, 3, 1, 1, bias=False)\n        self.comb_iter_4_right = nn.MaxPool2d(3, stride=2, padding=1)\n\n    def forward(self, x_conv0, x_stem_0):\n        x_left = self.conv_1x1(x_stem_0)\n\n        x_relu = self.relu(x_conv0)\n        # path 1\n        x_path1 = self.path_1(x_relu)\n        # path 2\n        x_path2 = self.path_2(x_relu[:, :, 1:, 1:])\n        # final path\n        x_right = self.final_path_bn(torch.cat([x_path1, x_path2], 1))\n\n        return ReductionCellBranchCombine(self, x_left, x_right)\n\n\nclass ReductionCell(nn.Module):\n\n    def __init__(self, in_channels_left, out_channels_left, in_channels_right, out_channels_right):\n        super(ReductionCell, self).__init__()\n        self.conv_prev_1x1 = nn.Sequential()\n        self.conv_prev_1x1.add_module('relu', nn.ReLU())\n        self.conv_prev_1x1.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n        self.conv_prev_1x1.add_module('bn', nn.BatchNorm2d(out_channels_left, eps=0.001, momentum=0.1, affine=True))\n\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_right, out_channels_right, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels_right, eps=0.001, momentum=0.1, affine=True))\n\n        self.comb_iter_0_left = BranchSeparables(out_channels_right, out_channels_right, 5, 2, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparables(out_channels_right, out_channels_right, 7, 2, 3, bias=False)\n\n        self.comb_iter_1_left = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_1_right = BranchSeparables(out_channels_right, out_channels_right, 7, 2, 3, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=2, padding=1, count_include_pad=False)\n        self.comb_iter_2_right = BranchSeparables(out_channels_right, out_channels_right, 5, 2, 2, bias=False)\n\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n        self.comb_iter_4_right = nn.MaxPool2d(3, stride=2, padding=1)\n\n    def forward(self, x, x_prev):\n        x_left = self.conv_1x1(x)\n        x_right = self.conv_prev_1x1(x_prev)\n        return ReductionCellBranchCombine(self, x_left, x_right)\n\n\ndef NormalCellBranchCombine(cell, x_left, x_right):\n    x_comb_iter_0_left = cell.comb_iter_0_left(x_right)\n    x_comb_iter_0_right = cell.comb_iter_0_right(x_left)\n    x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n    x_comb_iter_1_left = cell.comb_iter_1_left(x_left)\n    x_comb_iter_1_right = cell.comb_iter_1_right(x_left)\n    x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n    x_comb_iter_2_left = cell.comb_iter_2_left(x_right)\n    x_comb_iter_2 = x_comb_iter_2_left + x_left\n\n    x_comb_iter_3_left = cell.comb_iter_3_left(x_left)\n    x_comb_iter_3_right = cell.comb_iter_3_right(x_left)\n    x_comb_iter_3 = x_comb_iter_3_left + x_comb_iter_3_right\n\n    x_comb_iter_4_left = cell.comb_iter_4_left(x_right)\n    x_comb_iter_4 = x_comb_iter_4_left + x_right\n\n    x_out = torch.cat([x_left, x_comb_iter_0, x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n    return x_out\n\n\nclass FirstCell(nn.Module):\n\n    def __init__(self, in_channels_left, out_channels_left, in_channels_right, out_channels_right):\n        super(FirstCell, self).__init__()\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_right, out_channels_right, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels_right, eps=0.001, momentum=0.1, affine=True))\n\n        self.relu = nn.ReLU()\n        self.path_1 = nn.Sequential()\n        self.path_1.add_module('avgpool', nn.AvgPool2d(1, stride=2, count_include_pad=False))\n        self.path_1.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n        self.path_2 = nn.Sequential()\n        self.path_2.add_module('avgpool', nn.AvgPool2d(1, stride=2, ceil_mode=True, count_include_pad=False))\n        self.path_2.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n\n        self.final_path_bn = nn.BatchNorm2d(out_channels_left * 2, eps=0.001, momentum=0.1, affine=True)\n\n        self.comb_iter_0_left = BranchSeparables(out_channels_right, out_channels_right, 5, 1, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n        self.comb_iter_1_left = BranchSeparables(out_channels_right, out_channels_right, 5, 1, 2, bias=False)\n        self.comb_iter_1_right = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_3_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n    def forward(self, x, x_prev):\n        x_relu = self.relu(x_prev)\n        # path 1\n        x_path1 = self.path_1(x_relu)\n        # path 2\n        x_path2 = self.path_2(x_relu[:, :, 1:, 1:])\n        # final path\n        x_left = self.final_path_bn(torch.cat([x_path1, x_path2], 1))\n\n        x_right = self.conv_1x1(x)\n        \n        return NormalCellBranchCombine(self, x_left, x_right)\n        \n\nclass NormalCell(nn.Module):\n\n    def __init__(self, in_channels_left, out_channels_left, in_channels_right, out_channels_right):\n        super(NormalCell, self).__init__()\n        self.conv_prev_1x1 = nn.Sequential()\n        self.conv_prev_1x1.add_module('relu', nn.ReLU())\n        self.conv_prev_1x1.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n        self.conv_prev_1x1.add_module('bn', nn.BatchNorm2d(out_channels_left, eps=0.001, momentum=0.1, affine=True))\n\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_right, out_channels_right, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels_right, eps=0.001, momentum=0.1, affine=True))\n\n        self.comb_iter_0_left = BranchSeparables(out_channels_right, out_channels_right, 5, 1, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparables(out_channels_left, out_channels_left, 3, 1, 1, bias=False)\n\n        self.comb_iter_1_left = BranchSeparables(out_channels_left, out_channels_left, 5, 1, 2, bias=False)\n        self.comb_iter_1_right = BranchSeparables(out_channels_left, out_channels_left, 3, 1, 1, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_3_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n    def forward(self, x, x_prev):\n        x_left = self.conv_prev_1x1(x_prev)\n        x_right = self.conv_1x1(x)\n\n        return NormalCellBranchCombine(self, x_left, x_right)\n\n","execution_count":2,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## NASNet Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.autograd import Variable\nfrom collections import OrderedDict\n\n\nclass NASNet(nn.Module):\n    def __init__(self, num_stem_features, num_normal_cells, filters, scaling, skip_reduction, use_aux=True,\n                 num_classes=1000):\n        super(NASNet, self).__init__()\n        self.num_normal_cells = num_normal_cells\n        self.skip_reduction = skip_reduction\n        self.use_aux = use_aux\n        self.num_classes = num_classes\n\n        self.conv0 = nn.Sequential(OrderedDict([\n            ('conv', nn.Conv2d(3, num_stem_features, kernel_size=3, stride=2, bias=False)),\n            ('bn', nn.BatchNorm2d(num_stem_features, eps=0.001, momentum=0.1, affine=True))\n        ]))\n\n        self.cell_stem_0 = CellStem0(in_channels=num_stem_features,\n                                     out_channels=int(filters * scaling ** (-2)))\n        self.cell_stem_1 = CellStem1(in_channels_x=int(4 * filters * scaling ** (-2)),\n                                     in_channels_h=num_stem_features,\n                                     out_channels=int(filters * scaling ** (-1)))\n\n        x_channels = int(4 * filters * scaling ** (-1))\n        h_channels = int(4 * filters * scaling ** (-2))\n        cell_id = 0\n        branch_out_channels = filters\n        for i in range(3):\n            self.add_module('cell_{:d}'.format(cell_id), FirstCell(\n                in_channels_left=h_channels, out_channels_left=branch_out_channels // 2, in_channels_right=x_channels,\n                out_channels_right=branch_out_channels))\n            cell_id += 1\n            h_channels = x_channels\n            x_channels = 6 * branch_out_channels  # normal: concat 6 branches\n            for _ in range(num_normal_cells - 1):\n                self.add_module('cell_{:d}'.format(cell_id), NormalCell(\n                    in_channels_left=h_channels, out_channels_left=branch_out_channels, in_channels_right=x_channels,\n                    out_channels_right=branch_out_channels))\n                h_channels = x_channels\n                cell_id += 1\n            if i == 1 and self.use_aux:\n                self.aux_features = nn.Sequential(\n                    nn.ReLU(),\n                    nn.AvgPool2d(kernel_size=(5, 5), stride=(3, 3),\n                                 padding=(2, 2), count_include_pad=False),\n                    nn.Conv2d(in_channels=x_channels, out_channels=128, kernel_size=1, bias=False),\n                    nn.BatchNorm2d(num_features=128, eps=0.001, momentum=0.1, affine=True),\n                    nn.ReLU(),\n                    nn.Conv2d(in_channels=128, out_channels=768,\n                              kernel_size=((14 + 2) // 3, (14 + 2) // 3), bias=False),\n                    nn.BatchNorm2d(num_features=768, eps=1e-3, momentum=0.1, affine=True),\n                    nn.ReLU()\n                )\n                self.aux_linear = nn.Linear(768, num_classes)\n            # scaling\n            branch_out_channels *= scaling\n            if i < 2:\n                self.add_module('reduction_cell_{:d}'.format(i), ReductionCell(\n                    in_channels_left=h_channels, out_channels_left=branch_out_channels,\n                    in_channels_right=x_channels, out_channels_right=branch_out_channels))\n                x_channels = 4 * branch_out_channels  # reduce: concat 4 branches\n\n        self.linear = nn.Linear(x_channels, self.num_classes)  # large: 4032; mobile: 1056\n\n        self.num_params = sum([param.numel() for param in self.parameters()])\n        if self.use_aux:\n            self.num_params -= sum([param.numel() for param in self.aux_features.parameters()])\n            self.num_params -= sum([param.numel() for param in self.aux_linear.parameters()])\n\n    def features(self, x):\n        x_conv0 = self.conv0(x)\n        x_stem_0 = self.cell_stem_0(x_conv0)\n        x_stem_1 = self.cell_stem_1(x_conv0, x_stem_0)\n        prev_x, x = x_stem_0, x_stem_1\n        cell_id = 0\n        for i in range(3):\n            for _ in range(self.num_normal_cells):\n                new_x = self._modules['cell_{:d}'.format(cell_id)](x, prev_x)\n                prev_x, x = x, new_x\n                cell_id += 1\n            if i == 1 and self.training and self.use_aux:\n                x_aux = self.aux_features(x)\n            if i < 2:\n                new_x = self._modules['reduction_cell_{:d}'.format(i)](x, prev_x)\n                prev_x = x if not self.skip_reduction else prev_x\n                x = new_x\n        if self.training and self.use_aux:\n            return [x, x_aux]\n        return [x]\n\n    def logits(self, features):\n        x = F.relu(features, inplace=False)\n        x = F.avg_pool2d(x, kernel_size=x.size(2)).view(x.size(0), -1)\n        x = F.dropout(x, p=0.5, training=self.training)\n        x = self.linear(x)\n        return x\n\n    def forward(self, x):\n        x = self.features(x)\n        output = self.logits(x[0])\n        if self.training and self.use_aux:\n            x_aux = x[1].view(x[1].size(0), -1)\n            aux_output = self.aux_linear(x_aux)\n            return [output, aux_output]\n        return [output]\n\n\ndef NASNetAMobile(num_classes=1103):\n    return NASNet(32, 4, 44, 2, skip_reduction=False, use_aux=True, num_classes=num_classes)\n\n\ndef NASNetALarge(num_classes=1000):\n    return NASNet(96, 6, 168, 2, skip_reduction=True, use_aux=True, num_classes=num_classes)\n","execution_count":3,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model=NASNetAMobile(1001)","execution_count":4,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model.linear","execution_count":5,"outputs":[{"output_type":"execute_result","execution_count":5,"data":{"text/plain":"Linear(in_features=1056, out_features=1001, bias=True)"},"metadata":{}}]},{"metadata":{"trusted":true},"cell_type":"code","source":"weights=torch.load('../input/pytorch-model-zoo/nasnetamobile-7e03cead.pth')","execution_count":6,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#model=model.load_state_dict(weights,strict=False)","execution_count":15,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import h5py","execution_count":16,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#!pip install pretrainedmodels\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.utils.model_zoo as model_zoo\nfrom torch.autograd import Variable\nimport numpy as np\n\npretrained_settings = {\n    'nasnetamobile': {\n        'imagenet': {\n            'url': 'https://github.com/veronikayurchuk/pretrained-models.pytorch/releases/download/v1.0/nasnetmobile-7e03cead.pth.tar',\n            'input_space': 'RGB',\n            'input_size': [3, 224, 224], # resize 256\n            'input_range': [0, 1],\n            'mean': [0.5, 0.5, 0.5],\n            'std': [0.5, 0.5, 0.5],\n            'num_classes': 1000\n        },\n        # 'imagenet+background': {\n        #     # 'url': 'http://data.lip6.fr/cadene/pretrainedmodels/nasnetalarge-a1897284.pth',\n        #     'input_space': 'RGB',\n        #     'input_size': [3, 224, 224], # resize 256\n        #     'input_range': [0, 1],\n        #     'mean': [0.5, 0.5, 0.5],\n        #     'std': [0.5, 0.5, 0.5],\n        #     'num_classes': 1001\n        # }\n    }\n}\n\n\nclass MaxPoolPad(nn.Module):\n\n    def __init__(self):\n        super(MaxPoolPad, self).__init__()\n        self.pad = nn.ZeroPad2d((1, 0, 1, 0))\n        self.pool = nn.MaxPool2d(3, stride=2, padding=1)\n\n    def forward(self, x):\n        x = self.pad(x)\n        x = self.pool(x)\n        x = x[:, :, 1:, 1:].contiguous()\n        return x\n\n\nclass AvgPoolPad(nn.Module):\n\n    def __init__(self, stride=2, padding=1):\n        super(AvgPoolPad, self).__init__()\n        self.pad = nn.ZeroPad2d((1, 0, 1, 0))\n        self.pool = nn.AvgPool2d(3, stride=stride, padding=padding, count_include_pad=False)\n\n    def forward(self, x):\n        x = self.pad(x)\n        x = self.pool(x)\n        x = x[:, :, 1:, 1:].contiguous()\n        return x\n\n\nclass SeparableConv2d(nn.Module):\n\n    def __init__(self, in_channels, out_channels, dw_kernel, dw_stride, dw_padding, bias=False):\n        super(SeparableConv2d, self).__init__()\n        self.depthwise_conv2d = nn.Conv2d(in_channels, in_channels, dw_kernel,\n                                          stride=dw_stride,\n                                          padding=dw_padding,\n                                          bias=bias,\n                                          groups=in_channels)\n        self.pointwise_conv2d = nn.Conv2d(in_channels, out_channels, 1, stride=1, bias=bias)\n\n    def forward(self, x):\n        x = self.depthwise_conv2d(x)\n        x = self.pointwise_conv2d(x)\n        return x\n\n\nclass BranchSeparables(nn.Module):\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, name=None, bias=False):\n        super(BranchSeparables, self).__init__()\n        self.relu = nn.ReLU()\n        self.separable_1 = SeparableConv2d(in_channels, in_channels, kernel_size, stride, padding, bias=bias)\n        self.bn_sep_1 = nn.BatchNorm2d(in_channels, eps=0.001, momentum=0.1, affine=True)\n        self.relu1 = nn.ReLU()\n        self.separable_2 = SeparableConv2d(in_channels, out_channels, kernel_size, 1, padding, bias=bias)\n        self.bn_sep_2 = nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True)\n        self.name = name\n\n    def forward(self, x):\n        x = self.relu(x)\n        if self.name == 'specific':\n            x = nn.ZeroPad2d((1, 0, 1, 0))(x)\n        x = self.separable_1(x)\n        if self.name == 'specific':\n            x = x[:, :, 1:, 1:].contiguous()\n\n        x = self.bn_sep_1(x)\n        x = self.relu1(x)\n        x = self.separable_2(x)\n        x = self.bn_sep_2(x)\n        return x\n\n\nclass BranchSeparablesStem(nn.Module):\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, bias=False):\n        super(BranchSeparablesStem, self).__init__()\n        self.relu = nn.ReLU()\n        self.separable_1 = SeparableConv2d(in_channels, out_channels, kernel_size, stride, padding, bias=bias)\n        self.bn_sep_1 = nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True)\n        self.relu1 = nn.ReLU()\n        self.separable_2 = SeparableConv2d(out_channels, out_channels, kernel_size, 1, padding, bias=bias)\n        self.bn_sep_2 = nn.BatchNorm2d(out_channels, eps=0.001, momentum=0.1, affine=True)\n\n    def forward(self, x):\n        x = self.relu(x)\n        x = self.separable_1(x)\n        x = self.bn_sep_1(x)\n        x = self.relu1(x)\n        x = self.separable_2(x)\n        x = self.bn_sep_2(x)\n        return x\n\n\nclass BranchSeparablesReduction(BranchSeparables):\n\n    def __init__(self, in_channels, out_channels, kernel_size, stride, padding, z_padding=1, bias=False):\n        BranchSeparables.__init__(self, in_channels, out_channels, kernel_size, stride, padding, bias)\n        self.padding = nn.ZeroPad2d((z_padding, 0, z_padding, 0))\n\n    def forward(self, x):\n        x = self.relu(x)\n        x = self.padding(x)\n        x = self.separable_1(x)\n        x = x[:, :, 1:, 1:].contiguous()\n        x = self.bn_sep_1(x)\n        x = self.relu1(x)\n        x = self.separable_2(x)\n        x = self.bn_sep_2(x)\n        return x\n\n\nclass CellStem0(nn.Module):\n    def __init__(self, stem_filters, num_filters=42):\n        super(CellStem0, self).__init__()\n        self.num_filters = num_filters\n        self.stem_filters = stem_filters\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(self.stem_filters, self.num_filters, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(self.num_filters, eps=0.001, momentum=0.1, affine=True))\n\n        self.comb_iter_0_left = BranchSeparables(self.num_filters, self.num_filters, 5, 2, 2)\n        self.comb_iter_0_right = BranchSeparablesStem(self.stem_filters, self.num_filters, 7, 2, 3, bias=False)\n\n        self.comb_iter_1_left = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_1_right = BranchSeparablesStem(self.stem_filters, self.num_filters, 7, 2, 3, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=2, padding=1, count_include_pad=False)\n        self.comb_iter_2_right = BranchSeparablesStem(self.stem_filters, self.num_filters, 5, 2, 2, bias=False)\n\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(self.num_filters, self.num_filters, 3, 1, 1, bias=False)\n        self.comb_iter_4_right = nn.MaxPool2d(3, stride=2, padding=1)\n\n    def forward(self, x):\n        x1 = self.conv_1x1(x)\n\n        x_comb_iter_0_left = self.comb_iter_0_left(x1)\n        x_comb_iter_0_right = self.comb_iter_0_right(x)\n        x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n        x_comb_iter_1_left = self.comb_iter_1_left(x1)\n        x_comb_iter_1_right = self.comb_iter_1_right(x)\n        x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n        x_comb_iter_2_left = self.comb_iter_2_left(x1)\n        x_comb_iter_2_right = self.comb_iter_2_right(x)\n        x_comb_iter_2 = x_comb_iter_2_left + x_comb_iter_2_right\n\n        x_comb_iter_3_right = self.comb_iter_3_right(x_comb_iter_0)\n        x_comb_iter_3 = x_comb_iter_3_right + x_comb_iter_1\n\n        x_comb_iter_4_left = self.comb_iter_4_left(x_comb_iter_0)\n        x_comb_iter_4_right = self.comb_iter_4_right(x1)\n        x_comb_iter_4 = x_comb_iter_4_left + x_comb_iter_4_right\n\n        x_out = torch.cat([x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n        return x_out\n\n\nclass CellStem1(nn.Module):\n\n    def __init__(self, stem_filters, num_filters):\n        super(CellStem1, self).__init__()\n        self.num_filters = num_filters\n        self.stem_filters = stem_filters\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(2*self.num_filters, self.num_filters, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(self.num_filters, eps=0.001, momentum=0.1, affine=True))\n\n        self.relu = nn.ReLU()\n        self.path_1 = nn.Sequential()\n        self.path_1.add_module('avgpool', nn.AvgPool2d(1, stride=2, count_include_pad=False))\n        self.path_1.add_module('conv', nn.Conv2d(self.stem_filters, self.num_filters//2, 1, stride=1, bias=False))\n        self.path_2 = nn.ModuleList()\n        self.path_2.add_module('pad', nn.ZeroPad2d((0, 1, 0, 1)))\n        self.path_2.add_module('avgpool', nn.AvgPool2d(1, stride=2, count_include_pad=False))\n        self.path_2.add_module('conv', nn.Conv2d(self.stem_filters, self.num_filters//2, 1, stride=1, bias=False))\n\n        self.final_path_bn = nn.BatchNorm2d(self.num_filters, eps=0.001, momentum=0.1, affine=True)\n\n        self.comb_iter_0_left = BranchSeparables(self.num_filters, self.num_filters, 5, 2, 2, name='specific', bias=False)\n        self.comb_iter_0_right = BranchSeparables(self.num_filters, self.num_filters, 7, 2, 3, name='specific', bias=False)\n\n        # self.comb_iter_1_left = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_1_left = MaxPoolPad()\n        self.comb_iter_1_right = BranchSeparables(self.num_filters, self.num_filters, 7, 2, 3, name='specific', bias=False)\n\n        # self.comb_iter_2_left = nn.AvgPool2d(3, stride=2, padding=1, count_include_pad=False)\n        self.comb_iter_2_left = AvgPoolPad()\n        self.comb_iter_2_right = BranchSeparables(self.num_filters, self.num_filters, 5, 2, 2, name='specific', bias=False)\n\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(self.num_filters, self.num_filters, 3, 1, 1, name='specific', bias=False)\n        # self.comb_iter_4_right = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_4_right = MaxPoolPad()\n\n    def forward(self, x_conv0, x_stem_0):\n        x_left = self.conv_1x1(x_stem_0)\n\n        x_relu = self.relu(x_conv0)\n        # path 1\n        x_path1 = self.path_1(x_relu)\n        # path 2\n        x_path2 = self.path_2.pad(x_relu)\n        x_path2 = x_path2[:, :, 1:, 1:]\n        x_path2 = self.path_2.avgpool(x_path2)\n        x_path2 = self.path_2.conv(x_path2)\n        # final path\n        x_right = self.final_path_bn(torch.cat([x_path1, x_path2], 1))\n\n        x_comb_iter_0_left = self.comb_iter_0_left(x_left)\n        x_comb_iter_0_right = self.comb_iter_0_right(x_right)\n        x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n        x_comb_iter_1_left = self.comb_iter_1_left(x_left)\n        x_comb_iter_1_right = self.comb_iter_1_right(x_right)\n        x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n        x_comb_iter_2_left = self.comb_iter_2_left(x_left)\n        x_comb_iter_2_right = self.comb_iter_2_right(x_right)\n        x_comb_iter_2 = x_comb_iter_2_left + x_comb_iter_2_right\n\n        x_comb_iter_3_right = self.comb_iter_3_right(x_comb_iter_0)\n        x_comb_iter_3 = x_comb_iter_3_right + x_comb_iter_1\n\n        x_comb_iter_4_left = self.comb_iter_4_left(x_comb_iter_0)\n        x_comb_iter_4_right = self.comb_iter_4_right(x_left)\n        x_comb_iter_4 = x_comb_iter_4_left + x_comb_iter_4_right\n\n        x_out = torch.cat([x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n        return x_out\n\n\nclass FirstCell(nn.Module):\n\n    def __init__(self, in_channels_left, out_channels_left, in_channels_right, out_channels_right):\n        super(FirstCell, self).__init__()\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_right, out_channels_right, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels_right, eps=0.001, momentum=0.1, affine=True))\n\n        self.relu = nn.ReLU()\n        self.path_1 = nn.Sequential()\n        self.path_1.add_module('avgpool', nn.AvgPool2d(1, stride=2, count_include_pad=False))\n        self.path_1.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n        self.path_2 = nn.ModuleList()\n        self.path_2.add_module('pad', nn.ZeroPad2d((0, 1, 0, 1)))\n        self.path_2.add_module('avgpool', nn.AvgPool2d(1, stride=2, count_include_pad=False))\n        self.path_2.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n\n        self.final_path_bn = nn.BatchNorm2d(out_channels_left * 2, eps=0.001, momentum=0.1, affine=True)\n\n        self.comb_iter_0_left = BranchSeparables(out_channels_right, out_channels_right, 5, 1, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n        self.comb_iter_1_left = BranchSeparables(out_channels_right, out_channels_right, 5, 1, 2, bias=False)\n        self.comb_iter_1_right = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_3_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n    def forward(self, x, x_prev):\n        x_relu = self.relu(x_prev)\n        # path 1\n        x_path1 = self.path_1(x_relu)\n        # path 2\n        x_path2 = self.path_2.pad(x_relu)\n        x_path2 = x_path2[:, :, 1:, 1:]\n        x_path2 = self.path_2.avgpool(x_path2)\n        x_path2 = self.path_2.conv(x_path2)\n        # final path\n        x_left = self.final_path_bn(torch.cat([x_path1, x_path2], 1))\n\n        x_right = self.conv_1x1(x)\n\n        x_comb_iter_0_left = self.comb_iter_0_left(x_right)\n        x_comb_iter_0_right = self.comb_iter_0_right(x_left)\n        x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n        x_comb_iter_1_left = self.comb_iter_1_left(x_left)\n        x_comb_iter_1_right = self.comb_iter_1_right(x_left)\n        x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n        x_comb_iter_2_left = self.comb_iter_2_left(x_right)\n        x_comb_iter_2 = x_comb_iter_2_left + x_left\n\n        x_comb_iter_3_left = self.comb_iter_3_left(x_left)\n        x_comb_iter_3_right = self.comb_iter_3_right(x_left)\n        x_comb_iter_3 = x_comb_iter_3_left + x_comb_iter_3_right\n\n        x_comb_iter_4_left = self.comb_iter_4_left(x_right)\n        x_comb_iter_4 = x_comb_iter_4_left + x_right\n\n        x_out = torch.cat([x_left, x_comb_iter_0, x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n        return x_out\n\n\nclass NormalCell(nn.Module):\n\n    def __init__(self, in_channels_left, out_channels_left, in_channels_right, out_channels_right):\n        super(NormalCell, self).__init__()\n        self.conv_prev_1x1 = nn.Sequential()\n        self.conv_prev_1x1.add_module('relu', nn.ReLU())\n        self.conv_prev_1x1.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n        self.conv_prev_1x1.add_module('bn', nn.BatchNorm2d(out_channels_left, eps=0.001, momentum=0.1, affine=True))\n\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_right, out_channels_right, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels_right, eps=0.001, momentum=0.1, affine=True))\n\n        self.comb_iter_0_left = BranchSeparables(out_channels_right, out_channels_right, 5, 1, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparables(out_channels_left, out_channels_left, 3, 1, 1, bias=False)\n\n        self.comb_iter_1_left = BranchSeparables(out_channels_left, out_channels_left, 5, 1, 2, bias=False)\n        self.comb_iter_1_right = BranchSeparables(out_channels_left, out_channels_left, 3, 1, 1, bias=False)\n\n        self.comb_iter_2_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_3_left = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n\n    def forward(self, x, x_prev):\n        x_left = self.conv_prev_1x1(x_prev)\n        x_right = self.conv_1x1(x)\n\n        x_comb_iter_0_left = self.comb_iter_0_left(x_right)\n        x_comb_iter_0_right = self.comb_iter_0_right(x_left)\n        x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n        x_comb_iter_1_left = self.comb_iter_1_left(x_left)\n        x_comb_iter_1_right = self.comb_iter_1_right(x_left)\n        x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n        x_comb_iter_2_left = self.comb_iter_2_left(x_right)\n        x_comb_iter_2 = x_comb_iter_2_left + x_left\n\n        x_comb_iter_3_left = self.comb_iter_3_left(x_left)\n        x_comb_iter_3_right = self.comb_iter_3_right(x_left)\n        x_comb_iter_3 = x_comb_iter_3_left + x_comb_iter_3_right\n\n        x_comb_iter_4_left = self.comb_iter_4_left(x_right)\n        x_comb_iter_4 = x_comb_iter_4_left + x_right\n\n        x_out = torch.cat([x_left, x_comb_iter_0, x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n        return x_out\n\n\nclass ReductionCell0(nn.Module):\n\n    def __init__(self, in_channels_left, out_channels_left, in_channels_right, out_channels_right):\n        super(ReductionCell0, self).__init__()\n        self.conv_prev_1x1 = nn.Sequential()\n        self.conv_prev_1x1.add_module('relu', nn.ReLU())\n        self.conv_prev_1x1.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n        self.conv_prev_1x1.add_module('bn', nn.BatchNorm2d(out_channels_left, eps=0.001, momentum=0.1, affine=True))\n\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_right, out_channels_right, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels_right, eps=0.001, momentum=0.1, affine=True))\n\n        self.comb_iter_0_left = BranchSeparablesReduction(out_channels_right, out_channels_right, 5, 2, 2, bias=False)\n        self.comb_iter_0_right = BranchSeparablesReduction(out_channels_right, out_channels_right, 7, 2, 3, bias=False)\n\n        self.comb_iter_1_left = MaxPoolPad()\n        self.comb_iter_1_right = BranchSeparablesReduction(out_channels_right, out_channels_right, 7, 2, 3, bias=False)\n\n        self.comb_iter_2_left = AvgPoolPad()\n        self.comb_iter_2_right = BranchSeparablesReduction(out_channels_right, out_channels_right, 5, 2, 2, bias=False)\n\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparablesReduction(out_channels_right, out_channels_right, 3, 1, 1, bias=False)\n        self.comb_iter_4_right = MaxPoolPad()\n\n    def forward(self, x, x_prev):\n        x_left = self.conv_prev_1x1(x_prev)\n        x_right = self.conv_1x1(x)\n\n        x_comb_iter_0_left = self.comb_iter_0_left(x_right)\n        x_comb_iter_0_right = self.comb_iter_0_right(x_left)\n        x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n        x_comb_iter_1_left = self.comb_iter_1_left(x_right)\n        x_comb_iter_1_right = self.comb_iter_1_right(x_left)\n        x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n        x_comb_iter_2_left = self.comb_iter_2_left(x_right)\n        x_comb_iter_2_right = self.comb_iter_2_right(x_left)\n        x_comb_iter_2 = x_comb_iter_2_left + x_comb_iter_2_right\n\n        x_comb_iter_3_right = self.comb_iter_3_right(x_comb_iter_0)\n        x_comb_iter_3 = x_comb_iter_3_right + x_comb_iter_1\n\n        x_comb_iter_4_left = self.comb_iter_4_left(x_comb_iter_0)\n        x_comb_iter_4_right = self.comb_iter_4_right(x_right)\n        x_comb_iter_4 = x_comb_iter_4_left + x_comb_iter_4_right\n\n        x_out = torch.cat([x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n        return x_out\n\n\nclass ReductionCell1(nn.Module):\n\n    def __init__(self, in_channels_left, out_channels_left, in_channels_right, out_channels_right):\n        super(ReductionCell1, self).__init__()\n        self.conv_prev_1x1 = nn.Sequential()\n        self.conv_prev_1x1.add_module('relu', nn.ReLU())\n        self.conv_prev_1x1.add_module('conv', nn.Conv2d(in_channels_left, out_channels_left, 1, stride=1, bias=False))\n        self.conv_prev_1x1.add_module('bn', nn.BatchNorm2d(out_channels_left, eps=0.001, momentum=0.1, affine=True))\n\n        self.conv_1x1 = nn.Sequential()\n        self.conv_1x1.add_module('relu', nn.ReLU())\n        self.conv_1x1.add_module('conv', nn.Conv2d(in_channels_right, out_channels_right, 1, stride=1, bias=False))\n        self.conv_1x1.add_module('bn', nn.BatchNorm2d(out_channels_right, eps=0.001, momentum=0.1, affine=True))\n\n        self.comb_iter_0_left = BranchSeparables(out_channels_right, out_channels_right, 5, 2, 2, name='specific', bias=False)\n        self.comb_iter_0_right = BranchSeparables(out_channels_right, out_channels_right, 7, 2, 3, name='specific', bias=False)\n\n        # self.comb_iter_1_left = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_1_left = MaxPoolPad()\n        self.comb_iter_1_right = BranchSeparables(out_channels_right, out_channels_right, 7, 2, 3, name='specific', bias=False)\n\n        # self.comb_iter_2_left = nn.AvgPool2d(3, stride=2, padding=1, count_include_pad=False)\n        self.comb_iter_2_left = AvgPoolPad()\n        self.comb_iter_2_right = BranchSeparables(out_channels_right, out_channels_right, 5, 2, 2, name='specific', bias=False)\n\n        self.comb_iter_3_right = nn.AvgPool2d(3, stride=1, padding=1, count_include_pad=False)\n\n        self.comb_iter_4_left = BranchSeparables(out_channels_right, out_channels_right, 3, 1, 1, name='specific', bias=False)\n        # self.comb_iter_4_right = nn.MaxPool2d(3, stride=2, padding=1)\n        self.comb_iter_4_right =MaxPoolPad()\n\n    def forward(self, x, x_prev):\n        x_left = self.conv_prev_1x1(x_prev)\n        x_right = self.conv_1x1(x)\n\n        x_comb_iter_0_left = self.comb_iter_0_left(x_right)\n        x_comb_iter_0_right = self.comb_iter_0_right(x_left)\n        x_comb_iter_0 = x_comb_iter_0_left + x_comb_iter_0_right\n\n        x_comb_iter_1_left = self.comb_iter_1_left(x_right)\n        x_comb_iter_1_right = self.comb_iter_1_right(x_left)\n        x_comb_iter_1 = x_comb_iter_1_left + x_comb_iter_1_right\n\n        x_comb_iter_2_left = self.comb_iter_2_left(x_right)\n        x_comb_iter_2_right = self.comb_iter_2_right(x_left)\n        x_comb_iter_2 = x_comb_iter_2_left + x_comb_iter_2_right\n\n        x_comb_iter_3_right = self.comb_iter_3_right(x_comb_iter_0)\n        x_comb_iter_3 = x_comb_iter_3_right + x_comb_iter_1\n\n        x_comb_iter_4_left = self.comb_iter_4_left(x_comb_iter_0)\n        x_comb_iter_4_right = self.comb_iter_4_right(x_right)\n        x_comb_iter_4 = x_comb_iter_4_left + x_comb_iter_4_right\n\n        x_out = torch.cat([x_comb_iter_1, x_comb_iter_2, x_comb_iter_3, x_comb_iter_4], 1)\n        return x_out\n\n\nclass NASNetAMobile(nn.Module):\n    \"\"\"NASNetAMobile (4 @ 1056) \"\"\"\n\n    def __init__(self, num_classes=1001, stem_filters=32, penultimate_filters=1056, filters_multiplier=2):\n        super(NASNetAMobile, self).__init__()\n        self.num_classes = num_classes\n        self.stem_filters = stem_filters\n        self.penultimate_filters = penultimate_filters\n        self.filters_multiplier = filters_multiplier\n\n        filters = self.penultimate_filters // 24\n        # 24 is default value for the architecture\n\n        self.conv0 = nn.Sequential()\n        self.conv0.add_module('conv', nn.Conv2d(in_channels=3, out_channels=self.stem_filters, kernel_size=3, padding=0, stride=2,\n                                                bias=False))\n        self.conv0.add_module('bn', nn.BatchNorm2d(self.stem_filters, eps=0.001, momentum=0.1, affine=True))\n\n        self.cell_stem_0 = CellStem0(self.stem_filters, num_filters=filters // (filters_multiplier ** 2))\n        self.cell_stem_1 = CellStem1(self.stem_filters, num_filters=filters // filters_multiplier)\n\n        self.cell_0 = FirstCell(in_channels_left=filters, out_channels_left=filters//2, # 1, 0.5\n                                in_channels_right=2*filters, out_channels_right=filters) # 2, 1\n        self.cell_1 = NormalCell(in_channels_left=2*filters, out_channels_left=filters, # 2, 1\n                                 in_channels_right=6*filters, out_channels_right=filters) # 6, 1\n        self.cell_2 = NormalCell(in_channels_left=6*filters, out_channels_left=filters, # 6, 1\n                                 in_channels_right=6*filters, out_channels_right=filters) # 6, 1\n        self.cell_3 = NormalCell(in_channels_left=6*filters, out_channels_left=filters, # 6, 1\n                                 in_channels_right=6*filters, out_channels_right=filters) # 6, 1\n\n        self.reduction_cell_0 = ReductionCell0(in_channels_left=6*filters, out_channels_left=2*filters, # 6, 2\n                                               in_channels_right=6*filters, out_channels_right=2*filters) # 6, 2\n\n        self.cell_6 = FirstCell(in_channels_left=6*filters, out_channels_left=filters, # 6, 1\n                                in_channels_right=8*filters, out_channels_right=2*filters) # 8, 2\n        self.cell_7 = NormalCell(in_channels_left=8*filters, out_channels_left=2*filters, # 8, 2\n                                 in_channels_right=12*filters, out_channels_right=2*filters) # 12, 2\n        self.cell_8 = NormalCell(in_channels_left=12*filters, out_channels_left=2*filters, # 12, 2\n                                 in_channels_right=12*filters, out_channels_right=2*filters) # 12, 2\n        self.cell_9 = NormalCell(in_channels_left=12*filters, out_channels_left=2*filters, # 12, 2\n                                 in_channels_right=12*filters, out_channels_right=2*filters) # 12, 2\n\n        self.reduction_cell_1 = ReductionCell1(in_channels_left=12*filters, out_channels_left=4*filters, # 12, 4\n                                               in_channels_right=12*filters, out_channels_right=4*filters) # 12, 4\n\n        self.cell_12 = FirstCell(in_channels_left=12*filters, out_channels_left=2*filters, # 12, 2\n                                 in_channels_right=16*filters, out_channels_right=4*filters) # 16, 4\n        self.cell_13 = NormalCell(in_channels_left=16*filters, out_channels_left=4*filters, # 16, 4\n                                  in_channels_right=24*filters, out_channels_right=4*filters) # 24, 4\n        self.cell_14 = NormalCell(in_channels_left=24*filters, out_channels_left=4*filters, # 24, 4\n                                  in_channels_right=24*filters, out_channels_right=4*filters) # 24, 4\n        self.cell_15 = NormalCell(in_channels_left=24*filters, out_channels_left=4*filters, # 24, 4\n                                  in_channels_right=24*filters, out_channels_right=4*filters) # 24, 4\n\n        self.relu = nn.ReLU()\n        self.avg_pool = nn.AvgPool2d(7, stride=1, padding=0)\n        self.dropout = nn.Dropout()\n        self.last_linear = nn.Linear(24*filters, self.num_classes)\n\n    def features(self, input):\n        x_conv0 = self.conv0(input)\n        x_stem_0 = self.cell_stem_0(x_conv0)\n        x_stem_1 = self.cell_stem_1(x_conv0, x_stem_0)\n\n        x_cell_0 = self.cell_0(x_stem_1, x_stem_0)\n        x_cell_1 = self.cell_1(x_cell_0, x_stem_1)\n        x_cell_2 = self.cell_2(x_cell_1, x_cell_0)\n        x_cell_3 = self.cell_3(x_cell_2, x_cell_1)\n\n        x_reduction_cell_0 = self.reduction_cell_0(x_cell_3, x_cell_2)\n\n        x_cell_6 = self.cell_6(x_reduction_cell_0, x_cell_3)\n        x_cell_7 = self.cell_7(x_cell_6, x_reduction_cell_0)\n        x_cell_8 = self.cell_8(x_cell_7, x_cell_6)\n        x_cell_9 = self.cell_9(x_cell_8, x_cell_7)\n\n        x_reduction_cell_1 = self.reduction_cell_1(x_cell_9, x_cell_8)\n\n        x_cell_12 = self.cell_12(x_reduction_cell_1, x_cell_9)\n        x_cell_13 = self.cell_13(x_cell_12, x_reduction_cell_1)\n        x_cell_14 = self.cell_14(x_cell_13, x_cell_12)\n        x_cell_15 = self.cell_15(x_cell_14, x_cell_13)\n        return x_cell_15\n\n    def logits(self, features):\n        x = self.relu(features)\n        x = self.avg_pool(x)\n        x = x.view(x.size(0), -1)\n        x = self.dropout(x)\n        x = self.last_linear(x)\n        x = nn.Softmax(1)(x)\n        return x\n\n    def forward(self, input):\n        x = self.features(input)\n        x = self.logits(x)\n        return x\n\n\ndef nasnetamobile(num_classes=1000, pretrained='imagenet'):\n    r\"\"\"NASNetALarge model architecture from the\n    `\"NASNet\" <https://arxiv.org/abs/1707.07012>`_ paper.\n    \"\"\"\n    if pretrained:\n        settings = pretrained_settings['nasnetamobile'][pretrained]\n        assert num_classes == settings['num_classes'], \\\n            \"num_classes should be {}, but is {}\".format(settings['num_classes'], num_classes)\n\n        # both 'imagenet'&'imagenet+background' are loaded from same parameters\n        model = NASNetAMobile(num_classes=num_classes)\n        model.load_state_dict(model_zoo.load_url(settings['url'], map_location=None))\n\n       # if pretrained == 'imagenet':\n       #     new_last_linear = nn.Linear(model.last_linear.in_features, 1000)\n       #     new_last_linear.weight.data = model.last_linear.weight.data[1:]\n       #     new_last_linear.bias.data = model.last_linear.bias.data[1:]\n       #     model.last_linear = new_last_linear\n\n        model.input_space = settings['input_space']\n        model.input_size = settings['input_size']\n        model.input_range = settings['input_range']\n\n        model.mean = settings['mean']\n        model.std = settings['std']\n    else:\n        settings = pretrained_settings['nasnetamobile']['imagenet']\n        model = NASNetAMobile(num_classes=num_classes)\n        model.input_space = settings['input_space']\n        model.input_size = settings['input_size']\n        model.input_range = settings['input_range']\n\n        model.mean = settings['mean']\n        model.std = settings['std']\n    return model","execution_count":19,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#from pretrainedmodels import  nasnetamobile,nasnetalarge\nbase_model=NASNetAMobile()\n#base_model.last_linear","execution_count":25,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#base_model.load_state_dict(weights,strict=False)","execution_count":26,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#base_model=nasnetamobile(num_classes=1000)\n","execution_count":27,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"base_model.last_linear","execution_count":28,"outputs":[{"output_type":"execute_result","execution_count":28,"data":{"text/plain":"Linear(in_features=1056, out_features=1001, bias=True)"},"metadata":{}}]},{"metadata":{"trusted":true},"cell_type":"code","source":"new_linear=nn.Linear(1000,1103)","execution_count":29,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"base_model.last_linear=nn.Linear(1056,1103)","execution_count":30,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for param in list(base_model.children())[:-1]:\n    param.requires_grad=False","execution_count":31,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import gc\nimport os\nimport sys\nimport time\nimport random\nimport logging\nimport datetime as dt\n\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.utils.data as data\nimport torch.nn.functional as F\nimport torchvision as vision\n\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nfrom pathlib import Path\nfrom PIL import Image\nfrom contextlib import contextmanager\n\nfrom joblib import Parallel, delayed\nfrom tqdm import tqdm\nfrom fastprogress import master_bar, progress_bar\n\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import fbeta_score\n\ntorch.multiprocessing.set_start_method(\"spawn\")","execution_count":32,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"@contextmanager\ndef timer(name=\"Main\", logger=None):\n    t0 = time.time()\n    yield\n    msg = f\"[{name}] done in {time.time() - t0} s\"\n    if logger is not None:\n        logger.info(msg)\n    else:\n        print(msg)\n        \n\ndef get_logger(name=\"Main\", tag=\"exp\", log_dir=\"log/\"):\n    log_path = Path(log_dir)\n    path = log_path / tag\n    path.mkdir(exist_ok=True, parents=True)\n\n    logger = logging.getLogger(name)\n    logger.setLevel(logging.INFO)\n\n    fh = logging.FileHandler(\n        path / (dt.datetime.now().strftime(\"%Y-%m-%d-%H-%M-%S\") + \".log\"))\n    sh = logging.StreamHandler(sys.stdout)\n    formatter = logging.Formatter(\n        \"%(asctime)s %(name)s %(levelname)s %(message)s\")\n\n    fh.setFormatter(formatter)\n    sh.setFormatter(formatter)\n    logger.addHandler(fh)\n    logger.addHandler(sh)\n    return logger\n\n\ndef seed_torch(seed=1029):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True","execution_count":33,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"logger = get_logger(name=\"Main\", tag=\"Pytorch-nasnet\")\n","execution_count":34,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"labels = pd.read_csv(\"../input/imet-2019-fgvc6/labels.csv\")\ntrain = pd.read_csv(\"../input/imet-2019-fgvc6/train.csv\")\nsample = pd.read_csv(\"../input/imet-2019-fgvc6/sample_submission.csv\")\ntrain.head()","execution_count":35,"outputs":[{"output_type":"execute_result","execution_count":35,"data":{"text/plain":"                 id        attribute_ids\n0  1000483014d91860          147 616 813\n1  1000fe2e667721fe       51 616 734 813\n2  1001614cb89646ee                  776\n3  10041eb49b297c08  51 671 698 813 1092\n4  100501c227f8beea  13 404 492 903 1093","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>id</th>\n      <th>attribute_ids</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>1000483014d91860</td>\n      <td>147 616 813</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>1000fe2e667721fe</td>\n      <td>51 616 734 813</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>1001614cb89646ee</td>\n      <td>776</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>10041eb49b297c08</td>\n      <td>51 671 698 813 1092</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>100501c227f8beea</td>\n      <td>13 404 492 903 1093</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"DATA_ROOT='../input/imet-2019-fgvc6/'\nfrom collections import defaultdict, Counter\nimport random\n\nimport pandas as pd\nimport tqdm\n\ndef make_folds(n_folds: int) -> pd.DataFrame:\n    df = pd.read_csv(DATA_ROOT+'train.csv')\n    cls_counts = Counter(cls for classes in df['attribute_ids'].str.split()\n                         for cls in classes)\n    fold_cls_counts = defaultdict(int)\n    folds = [-1] * len(df)\n    for item in tqdm.tqdm(df.sample(frac=1, random_state=42).itertuples(),\n                          total=len(df)):\n        cls = min(item.attribute_ids.split(), key=lambda cls: cls_counts[cls])\n        fold_counts = [(f, fold_cls_counts[f, cls]) for f in range(n_folds)]\n        min_count = min([count for _, count in fold_counts])\n        random.seed(item.Index)\n        fold = random.choice([f for f, count in fold_counts\n                              if count == min_count])\n        folds[item.Index] = fold\n        for cls in item.attribute_ids.split():\n            fold_cls_counts[fold, cls] += 1\n    df['fold'] = folds\n    return df\n\n","execution_count":36,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df = make_folds(n_folds=5)\ndf.to_csv('folds.csv', index=None)","execution_count":37,"outputs":[{"output_type":"stream","text":"100%|██████████| 109237/109237 [00:02<00:00, 48145.85it/s]\n","name":"stderr"}]},{"metadata":{"trusted":true},"cell_type":"code","source":"import random\nimport math\n\nfrom PIL import Image\nfrom torchvision.transforms import (\n    ToTensor, Normalize, Compose, Resize, CenterCrop, RandomCrop,\n    RandomHorizontalFlip)\n\n\nclass RandomSizedCrop:\n    \"\"\"Random crop the given PIL.Image to a random size\n    of the original size and and a random aspect ratio\n    of the original aspect ratio.\n    size: size of the smaller edge\n    interpolation: Default: PIL.Image.BILINEAR\n    \"\"\"\n\n    def __init__(self, size, interpolation=Image.BILINEAR,\n                 min_aspect=4/5, max_aspect=5/4,\n                 min_area=0.25, max_area=1):\n        self.size = size\n        self.interpolation = interpolation\n        self.min_aspect = min_aspect\n        self.max_aspect = max_aspect\n        self.min_area = min_area\n        self.max_area = max_area\n\n    def __call__(self, img):\n        for attempt in range(10):\n            area = img.size[0] * img.size[1]\n            target_area = random.uniform(self.min_area, self.max_area) * area\n            aspect_ratio = random.uniform(self.min_aspect, self.max_aspect)\n\n            w = int(round(math.sqrt(target_area * aspect_ratio)))\n            h = int(round(math.sqrt(target_area / aspect_ratio)))\n\n            if random.random() < 0.5:\n                w, h = h, w\n\n            if w <= img.size[0] and h <= img.size[1]:\n                x1 = random.randint(0, img.size[0] - w)\n                y1 = random.randint(0, img.size[1] - h)\n\n                img = img.crop((x1, y1, x1 + w, y1 + h))\n                assert(img.size == (w, h))\n\n                return img.resize((self.size, self.size), self.interpolation)\n\n        # Fallback\n        scale = Resize(self.size, interpolation=self.interpolation)\n        crop = CenterCrop(self.size)\n        return crop(scale(img))\n\n\ntrain_transform = Compose([\n    RandomCrop(288),\n    RandomHorizontalFlip(),\n])\n\n\ntest_transform = Compose([\n    RandomCrop(288),\n    RandomHorizontalFlip(),\n])\n\n\ntensor_transform = Compose([\n    ToTensor(),\n    Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])","execution_count":38,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import argparse\nfrom itertools import islice\nimport json\nfrom pathlib import Path\nimport shutil\nimport warnings\nfrom typing import Dict\n\nfrom torch.utils.data import DataLoader,Dataset\nimport numpy as np\nimport pandas as pd\nfrom sklearn.metrics import fbeta_score\nfrom sklearn.exceptions import UndefinedMetricWarning\nimport torch\nfrom torch import nn, cuda\nfrom torch.optim import Adam\nimport tqdm","execution_count":39,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"run_root = Path('model1')\nfolds = pd.read_csv('folds.csv')\ntrain_root = DATA_ROOT+'train'\ntrain_fold = folds[folds['fold'] != 4]\nvalid_fold = folds[folds['fold'] == 4]\n\ndef make_loader(df: pd.DataFrame, image_transform) -> DataLoader:\n    return DataLoader(\n        TrainDataset(train_root, df, image_transform),\n        shuffle=True,\n        batch_size=64,\n        num_workers=4,\n        pin_memory=True,\n    )\n","execution_count":40,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"use_cuda = torch.cuda.is_available()\n#fresh_params = list(model.fresh_params())\nall_params = list(base_model.parameters())\nif use_cuda:\n    base_model = base_model.cuda()\n","execution_count":41,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device= 'cuda' if torch.cuda.is_available() else 'cpu'","execution_count":43,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device","execution_count":44,"outputs":[{"output_type":"execute_result","execution_count":44,"data":{"text/plain":"'cpu'"},"metadata":{}}]},{"metadata":{"trusted":true},"cell_type":"code","source":"from torch.utils.data import  Dataset,DataLoader\nfrom pathlib import Path\nfrom typing import Callable, List\n\nimport cv2\nimport pandas as pd\nfrom PIL import Image\n\nN_CLASSES=1103\nclass TrainDataset(Dataset):\n    def __init__(self, root: Path, df: pd.DataFrame,\n                 image_transform: Callable):\n        super().__init__()\n        self._root = root\n        self._df = df\n        self._image_transform = image_transform\n        #self._debug = debug\n\n    def __len__(self):\n        return len(self._df)\n\n    def __getitem__(self, idx: int):\n        item = self._df.iloc[idx]\n        image = load_transform_image(\n            item, self._root, self._image_transform)\n        target = torch.zeros(N_CLASSES)\n        for cls in item.attribute_ids.split():\n            target[int(cls)] = 1\n        return image, target\n\n\nclass TTADataset:\n    def __init__(self, root: Path, df: pd.DataFrame,\n                 image_transform: Callable, tta: int):\n        self._root = root\n        self._df = df\n        self._image_transform = image_transform\n        self._tta = tta\n\n    def __len__(self):\n        return len(self._df) * self._tta\n\n    def __getitem__(self, idx):\n        item = self._df.iloc[idx % len(self._df)]\n        image = load_transform_image(item, self._root, self._image_transform)\n        return image, item.id\n\n\ndef load_transform_image(\n        item, root: Path, image_transform: Callable):\n    image = load_image(item, root)\n    image = image_transform(image)\n    \n    return tensor_transform(image)\n\nfrom PIL import Image\ndef load_image(item, root: Path) -> Image.Image:\n    image = Image.open(str(root + f'/{item.id}.png'))\n    #image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    return image#Image.fromarray(image)\n\n\ndef get_ids(root: Path) -> List[str]:\n    return sorted({p.name.split('_')[0] for p in root.glob('*.png')})","execution_count":46,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_loader = make_loader(train_fold, train_transform)\ndef make_val_loader(df: pd.DataFrame, image_transform) -> DataLoader:\n    return DataLoader(\n        TrainDataset(train_root, df, image_transform),\n        shuffle=False,\n        batch_size=64,\n        num_workers=4,\n        pin_memory=True,\n    )\n\nvalid_loader = make_val_loader(valid_fold, test_transform)\ncriterion = nn.BCEWithLogitsLoss(reduction='none')\ntrain_kwargs = dict(\n    model=model,\n    criterion=criterion,\n    train_loader=train_loader,\n    valid_loader=valid_loader,\n    patience=9,\n    init_optimizer=lambda params, lr: Adam(params, lr),\n    use_cuda=use_cuda,\n)","execution_count":47,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def binarize_prediction(probabilities, threshold: float, argsorted=None,\n                        min_labels=1, max_labels=10):\n    \"\"\" Return matrix of 0/1 predictions, same shape as probabilities.\n    \"\"\"\n    assert probabilities.shape[1] == N_CLASSES\n    if argsorted is None:\n        argsorted = probabilities.argsort(axis=1)\n    max_mask = _make_mask(argsorted, max_labels)\n    min_mask = _make_mask(argsorted, min_labels)\n    prob_mask = probabilities > threshold\n    return (max_mask & prob_mask) | min_mask\n\n\ndef _make_mask(argsorted, top_n: int):\n    mask = np.zeros_like(argsorted, dtype=np.uint8)\n    col_indices = argsorted[:, -top_n:].reshape(-1)\n    row_indices = [i // top_n for i in range(len(col_indices))]\n    mask[row_indices, col_indices] = 1\n    return mask\n\n\ndef _reduce_loss(loss):\n    return loss.sum() / loss.shape[0]\n\n","execution_count":48,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"tta=4\nfold=5","execution_count":49,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"n_epochs=10\nbatch_size=64\nlr=1e-4\nsoftMax = nn.Softmax()\nreport_each=1\nCE_loss = nn.CrossEntropyLoss()\nif torch.cuda.is_available():\n        criterion.cuda()\n        softMax.cuda()\n        CE_loss.cuda()\n        \noptimizerG = Adam(base_model.parameters(), lr=lr, betas=(0.9, 0.99), amsgrad=False)\n    \nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizerG, mode='max', patience=4, verbose=True,\n                                                   factor=10 ** -0.5)\n","execution_count":50,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"Losses = []\nsave = lambda ep: torch.save({\n        'model': model.state_dict(),\n        'epoch': ep,\n        'step': step,\n        'best_valid_loss': best_valid_loss\n    }, str(model_path))\nvalid_losses = []\ndef validation(\n        model: nn.Module, criterion, valid_loader, use_cuda,\n        ) -> Dict[str, float]:\n    model.eval()\n    all_losses, all_predictions, all_targets = [], [], []\n    with torch.no_grad():\n        for inputs, targets in valid_loader:\n            all_targets.append(targets.numpy().copy())\n            if use_cuda:\n                inputs, targets = inputs.to(device), targets.to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n            all_losses.append(_reduce_loss(loss).item())\n            predictions = torch.sigmoid(outputs)\n            all_predictions.append(predictions.cpu().numpy())\n    all_predictions = np.concatenate(all_predictions)\n    all_targets = np.concatenate(all_targets)\n\n    def get_score(y_pred):\n        with warnings.catch_warnings():\n            warnings.simplefilter('ignore', category=UndefinedMetricWarning)\n            return fbeta_score(\n                all_targets, y_pred, beta=2, average='samples')\n\n    metrics = {}\n    argsorted = all_predictions.argsort(axis=1)\n    for threshold in [0.05, 0.06, 0.07, 0.08, 0.09, 0.10, 0.11, 0.12, 0.13, 0.14, 0.15, 0.20]:\n        metrics[f'valid_f2_th_{threshold:.2f}'] = get_score(\n            binarize_prediction(all_predictions, threshold, argsorted))\n    metrics['valid_loss'] = np.mean(all_losses)\n    print(' | '.join(f'{k} {v:.3f}' for k, v in sorted(\n        metrics.items(), key=lambda kv: -kv[1])))\n\n    return metrics\n\nfor i in range(n_epochs):\n    base_model.train()\n    success = 0\n    totalImages = len(train_loader)\n    mean_loss=0\n    print (f'Epoch: {i} , Lr: {lr} ')\n    for j, data in enumerate(train_loader):\n        image,label=data\n        optimizerG.zero_grad()\n        image=Variable(image).to(device)\n        target=Variable(label).to(device)\n        base_model.zero_grad()\n        \n        prediction=base_model(image)\n        loss=_reduce_loss(criterion(prediction,target))\n        batch=image.size(0)\n        (batch*loss).backward()\n        optimizerG.step()\n        Losses.append(loss.item())\n        mean_loss=np.mean(losses[-report_each:])\n        print(f'{mean_loss: .3f}')\n        save(i + 1)\n        valid_metrics = validation(base_model, criterion, valid_loader, use_cuda)\n        valid_loss = valid_metrics['valid_loss']\n        valid_losses.append(valid_loss)\n        if valid_loss < best_valid_loss:\n            best_valid_loss = valid_loss\n            shutil.copy(str(model_path), str(best_model_path))\n        \n        ","execution_count":null,"outputs":[{"output_type":"stream","text":"Epoch: 0 , Lr: 0.0001 \n","name":"stdout"}]},{"metadata":{"trusted":true},"cell_type":"code","source":"#train(params=all_params, **train_kwargs)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.4","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}