{"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":"BATCH_SIZE = 8\nPROPOSAL_NUM = 6\nCAT_NUM = 4\nINPUT_SIZE = (448, 448)  # (w, h)\nLR = 0.001\nWD = 1e-4\nSAVE_FREQ = 1\n#resume = '../input/models4nts/030.ckpt'\nresume = ''\nroot='../input/sorghum-id-fgvc-9'\nIMG_TRAIN_VAL='../input/sorghum-id-fgvc-9/train_images'\nIMG_TEST='../input/sorghum-id-fgvc-9/test/'\nsubmission='../input/submissioncsv'\ntest_model = ''\n# save_dir = 'output/kaggle/working/models'\nsave_dir = 'models'\n","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:51:38.066432Z","iopub.execute_input":"2022-06-25T04:51:38.066699Z","iopub.status.idle":"2022-06-25T04:51:38.072358Z","shell.execute_reply.started":"2022-06-25T04:51:38.066670Z","shell.execute_reply":"2022-06-25T04:51:38.071124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport scipy.misc\nimport os\nfrom PIL import Image\nfrom torchvision import transforms\nimport imageio\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nimport matplotlib.pyplot as plt\nfrom PIL import Image # pil及cv库，用于图片的打开修改等工作\nimport torchvision\nimport cv2\nprint(torch.__version__)  # 检查pytorch的版本\nprint(torch.cuda.is_available())  # 确定你的电脑的cuda版本，需要与你安装的pytorch的cuda版本匹配\nprint(torch.version.cuda)  # 确定你的电脑的cuda版本，需要与安装的pytorchd版本匹配\n\n# from config import INPUT_SIZE","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:51:38.131468Z","iopub.execute_input":"2022-06-25T04:51:38.131654Z","iopub.status.idle":"2022-06-25T04:51:38.611507Z","shell.execute_reply.started":"2022-06-25T04:51:38.131631Z","shell.execute_reply":"2022-06-25T04:51:38.610735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_msg=pd.read_csv(os.path.join(root,'train_cultivar_mapping.csv'))\ntrain_msg.dropna(inplace=True)  #这行可以去掉标题和存在数据缺失的项\nprint (f\"number of train data = {len(train_msg)}\")\ntrain_msg.head()\nsorghum_type = list(train_msg[\"cultivar\"].unique())\nprint(sorghum_type)\nprint(len(sorghum_type))#100个种类\n# train_msg[\"dir\"] = train_msg[\"image\"].apply( lambda image:TRAIN_IMG+image )\n# train_msg[\"dir_is_true\"] = train_msg[\"dir\"].apply(lambda dir: os.path.exists(dir))\n# train_msg[\"sorghum_type\"] = train_msg[\"cultivar\"].map(lambda type:sorghum_type.index(type))\n# train_msg.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:51:38.613218Z","iopub.execute_input":"2022-06-25T04:51:38.613460Z","iopub.status.idle":"2022-06-25T04:51:38.671679Z","shell.execute_reply.started":"2022-06-25T04:51:38.613426Z","shell.execute_reply":"2022-06-25T04:51:38.671018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_msg[\"dir\"] = train_msg[\"image\"].apply( lambda image:IMG_TRAIN_VAL+'/'+image )\ntrain_msg[\"dir_is_true\"] = train_msg[\"dir\"].apply(lambda dir: os.path.exists(dir))\ntrain_msg[\"sorghum_type\"] = train_msg[\"cultivar\"].map(lambda type:sorghum_type.index(type))\ntrain_msg.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:51:38.673954Z","iopub.execute_input":"2022-06-25T04:51:38.674444Z","iopub.status.idle":"2022-06-25T04:52:08.082617Z","shell.execute_reply.started":"2022-06-25T04:51:38.674402Z","shell.execute_reply":"2022-06-25T04:52:08.081918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"RATIO = 0.8\ntfms = transforms.Compose([transforms.RandomAffine(30,scale=(1,1.2)),\n    transforms.RandomResizedCrop(512), \n    transforms.RandomAdjustSharpness(2.),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(90),\n    transforms.RandomApply(torch.nn.ModuleList([transforms.GaussianBlur(5)])),\n    transforms.ToTensor(),\n    transforms.Normalize([0.385, 0.356, 0.306], [0.229, 0.224, 0.225]),])\n# tfms=transforms.Compose([\n#     transforms.Resize((600, 600), Image.BILINEAR)(),\n#     transforms.RandomCrop(INPUT_SIZE)(),\n#     transforms.RandomHorizontalFlip()(),\n#     transforms.ToTensor(),\n#     transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n# ])\ntsfm = transforms.Compose([transforms.RandomResizedCrop(512), \n    transforms.ToTensor(),\n    transforms.Normalize([0.385, 0.356, 0.306], [0.229, 0.224, 0.225]),])\nclass MyDataset(Dataset):\n    def __init__(self,train_msg,mode = \"train\",trans = tfms):\n        \"\"\"\n        :param train_msg: 整理后的地址和品种信息\n        :param mode: 训练集/验证集\n        :param trans: 数据预处理函数\n        \"\"\"\n        self.trans = trans\n        self.mode = mode\n        num_all = int(len(train_msg))\n        num_train = int(num_all * RATIO);\n        num_val = num_all - num_train;\n        if self.mode == \"train\":\n            img_list = train_msg[:num_train]\n        elif self.mode == \"val\":\n            img_list = train_msg[num_train:num_all]\n        self.length = len(img_list)\n        self.label = img_list[\"sorghum_type\"].values\n        self.img_dir =  img_list[\"dir\"].values\n        self.img_true = img_list[\"dir_is_true\"].values\n    \n    def __len__(self):\n        return self.length\n    \n    def __getitem__(self,item ):\n        if self.trans is None:\n                trans = transforms.Compose([\n                    transforms.Resize([IMG_WIDE, IMG_HEIGHT]),\n                    transforms.ToTensor(),\n                ])\n        else:\n            trans = self.trans\n            \n        if self.img_true[item]:\n            label = torch.tensor(self.label[item], dtype=torch.long)\n            image_path = self.img_dir[item]\n            image=imageio.imread(image_path)\n            image = Image.fromarray(image, mode='RGB')\n            image = transforms.Resize((600, 600), Image.BILINEAR)(image)\n            image = transforms.RandomCrop(INPUT_SIZE)(image)\n            image = transforms.RandomHorizontalFlip()(image)\n            image = transforms.ToTensor()(image)\n            image = transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])(image)\n#             image = Image.open(image_path)\n#             image = trans(image)\n        else:\n            self.length -= 1\n            print(f\"loss picture {item}\")\n        return image,  label\n\nclass TestDataset(Dataset):\n    def __init__(self,test_msg,trans = tsfm):\n        self.trans = trans\n        self.length = len(test_msg)\n        self.name  =test_msg[\"filename\"].values\n        self.img_dir =  test_msg[\"dir\"].values\n        self.img_true = test_msg[\"dir_is_true\"].values\n        img_list=test_msg\n        self.label = img_list[\"sorghum_type\"].values\n        \n    def __len__(self):\n        return self.length\n    \n    def __getitem__(self,item ):\n        if self.trans is None:\n                trans = transforms.Compose([\n                    transforms.Resize([IMG_WIDE, IMG_HEIGHT]),\n                    transforms.ToTensor(),\n                ])\n        else:\n            trans = self.trans\n            \n        if self.img_true[item]:\n            label = torch.tensor(self.label[item], dtype=torch.long)\n            image_path = self.img_dir[item]\n            image=imageio.imread(image_path)\n            image = Image.fromarray(image, mode='RGB')\n            image = transforms.Resize((600, 600), Image.BILINEAR)(image)\n            image = transforms.RandomCrop(INPUT_SIZE)(image)\n            image = transforms.RandomHorizontalFlip()(image)\n            image = transforms.ToTensor()(image)\n            image = transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])(image)\n            \n#             image_path = self.img_dir[item]\n#             image = Image.open(image_path)\n#             image = trans(image)\n        else:\n            self.length -= 1\n            print(f\"loss picture {item}\")\n            \n        return image,label #self.name[item]","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:52:08.085647Z","iopub.execute_input":"2022-06-25T04:52:08.085889Z","iopub.status.idle":"2022-06-25T04:52:08.108270Z","shell.execute_reply.started":"2022-06-25T04:52:08.085860Z","shell.execute_reply":"2022-06-25T04:52:08.107514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"KMP_DUPLICATE_LIB_OK\"] = \"TRUE\"\ndataset_val = MyDataset(train_msg, mode='train')\nval_loader = DataLoader(dataset_val, batch_size=40, shuffle=True)\ni = 0\n#取验证集中的40张图片显示\nfor batch in val_loader:\n    \n    if (i == 0):\n        images,labels= batch\n        print(type(images))\n        i += 1\n    else:\n        break\ngrid = torchvision.utils.make_grid(images, nrow=10)\nprint(labels)\nplt.figure(figsize=(20, 20))\nplt.imshow(np.transpose(grid, (1, 2, 0)))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:52:08.109291Z","iopub.execute_input":"2022-06-25T04:52:08.110000Z","iopub.status.idle":"2022-06-25T04:52:15.924445Z","shell.execute_reply.started":"2022-06-25T04:52:08.109963Z","shell.execute_reply":"2022-06-25T04:52:15.923287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# anchors","metadata":{}},{"cell_type":"code","source":"_default_anchors_setting = (\n    dict(layer='p3', stride=32, size=48, scale=[2 ** (1. / 3.), 2 ** (2. / 3.)], aspect_ratio=[0.667, 1, 1.5]),\n    dict(layer='p4', stride=64, size=96, scale=[2 ** (1. / 3.), 2 ** (2. / 3.)], aspect_ratio=[0.667, 1, 1.5]),\n    dict(layer='p5', stride=128, size=192, scale=[1, 2 ** (1. / 3.), 2 ** (2. / 3.)], aspect_ratio=[0.667, 1, 1.5]),\n)\n\n\ndef generate_default_anchor_maps(anchors_setting=None, input_shape=INPUT_SIZE):\n    \"\"\"\n    generate default anchor\n\n    :param anchors_setting: all informations of anchors\n    :param input_shape: shape of input images, e.g. (h, w)\n    :return: center_anchors: # anchors * 4 (oy, ox, h, w)\n             edge_anchors: # anchors * 4 (y0, x0, y1, x1)\n             anchor_area: # anchors * 1 (area)\n    \"\"\"\n    if anchors_setting is None:\n        anchors_setting = _default_anchors_setting\n\n    center_anchors = np.zeros((0, 4), dtype=np.float32)\n    edge_anchors = np.zeros((0, 4), dtype=np.float32)\n    anchor_areas = np.zeros((0,), dtype=np.float32)\n    input_shape = np.array(input_shape, dtype=int)\n\n    for anchor_info in anchors_setting:\n\n        stride = anchor_info['stride']\n        size = anchor_info['size']\n        scales = anchor_info['scale']\n        aspect_ratios = anchor_info['aspect_ratio']\n\n        output_map_shape = np.ceil(input_shape.astype(np.float32) / stride)\n        output_map_shape = output_map_shape.astype(np.int)\n        output_shape = tuple(output_map_shape) + (4,)\n        ostart = stride / 2.\n        oy = np.arange(ostart, ostart + stride * output_shape[0], stride)\n        oy = oy.reshape(output_shape[0], 1)\n        ox = np.arange(ostart, ostart + stride * output_shape[1], stride)\n        ox = ox.reshape(1, output_shape[1])\n        center_anchor_map_template = np.zeros(output_shape, dtype=np.float32)\n        center_anchor_map_template[:, :, 0] = oy\n        center_anchor_map_template[:, :, 1] = ox\n        for scale in scales:\n            for aspect_ratio in aspect_ratios:\n                center_anchor_map = center_anchor_map_template.copy()\n                center_anchor_map[:, :, 2] = size * scale / float(aspect_ratio) ** 0.5\n                center_anchor_map[:, :, 3] = size * scale * float(aspect_ratio) ** 0.5\n\n                edge_anchor_map = np.concatenate((center_anchor_map[..., :2] - center_anchor_map[..., 2:4] / 2.,\n                                                  center_anchor_map[..., :2] + center_anchor_map[..., 2:4] / 2.),\n                                                 axis=-1)\n                anchor_area_map = center_anchor_map[..., 2] * center_anchor_map[..., 3]\n                center_anchors = np.concatenate((center_anchors, center_anchor_map.reshape(-1, 4)))\n                edge_anchors = np.concatenate((edge_anchors, edge_anchor_map.reshape(-1, 4)))\n                anchor_areas = np.concatenate((anchor_areas, anchor_area_map.reshape(-1)))\n\n    return center_anchors, edge_anchors, anchor_areas\n\n\ndef hard_nms(cdds, topn=10, iou_thresh=0.25):\n    if not (type(cdds).__module__ == 'numpy' and len(cdds.shape) == 2 and cdds.shape[1] >= 5):\n        raise TypeError('edge_box_map should be N * 5+ ndarray')\n\n    cdds = cdds.copy()\n    indices = np.argsort(cdds[:, 0])\n    cdds = cdds[indices]\n    cdd_results = []\n\n    res = cdds\n\n    while res.any():\n        cdd = res[-1]\n        cdd_results.append(cdd)\n        if len(cdd_results) == topn:\n            return np.array(cdd_results)\n        res = res[:-1]\n\n        start_max = np.maximum(res[:, 1:3], cdd[1:3])\n        end_min = np.minimum(res[:, 3:5], cdd[3:5])\n        lengths = end_min - start_max\n        intersec_map = lengths[:, 0] * lengths[:, 1]\n        intersec_map[np.logical_or(lengths[:, 0] < 0, lengths[:, 1] < 0)] = 0\n        iou_map_cur = intersec_map / ((res[:, 3] - res[:, 1]) * (res[:, 4] - res[:, 2]) + (cdd[3] - cdd[1]) * (\n            cdd[4] - cdd[2]) - intersec_map)\n        res = res[iou_map_cur < iou_thresh]\n\n    return np.array(cdd_results)\n","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:52:15.927245Z","iopub.execute_input":"2022-06-25T04:52:15.928664Z","iopub.status.idle":"2022-06-25T04:52:15.974043Z","shell.execute_reply.started":"2022-06-25T04:52:15.928627Z","shell.execute_reply":"2022-06-25T04:52:15.973313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# utils","metadata":{}},{"cell_type":"code","source":"from __future__ import print_function\nimport os\nimport sys\nimport time\nimport logging\n\n# _, term_width = os.popen('stty size', 'r').read().split()\n# term_width = int(term_width)\nterm_width = 80\n\nTOTAL_BAR_LENGTH = 40.\nlast_time = time.time()\nbegin_time = last_time\n\n\ndef progress_bar(current, total, msg=None):\n    global last_time, begin_time\n    if current == 0:\n        begin_time = time.time()  # Reset for new bar.\n\n    cur_len = int(TOTAL_BAR_LENGTH * current / total)\n    rest_len = int(TOTAL_BAR_LENGTH - cur_len) - 1\n\n    sys.stdout.write(' [')\n    for i in range(cur_len):\n        sys.stdout.write('=')\n    sys.stdout.write('>')\n    for i in range(rest_len):\n        sys.stdout.write('.')\n    sys.stdout.write(']')\n\n    cur_time = time.time()\n    step_time = cur_time - last_time\n    last_time = cur_time\n    tot_time = cur_time - begin_time\n\n    L = []\n    L.append('  Step: %s' % format_time(step_time))\n    L.append(' | Tot: %s' % format_time(tot_time))\n    if msg:\n        L.append(' | ' + msg)\n\n    msg = ''.join(L)\n    sys.stdout.write(msg)\n    for i in range(term_width - int(TOTAL_BAR_LENGTH) - len(msg) - 3):\n        sys.stdout.write(' ')\n\n    # Go back to the center of the bar.\n    for i in range(term_width - int(TOTAL_BAR_LENGTH / 2)):\n        sys.stdout.write('\\b')\n    sys.stdout.write(' %d/%d ' % (current + 1, total))\n\n    if current < total - 1:\n        sys.stdout.write('\\r')\n    else:\n        sys.stdout.write('\\n')\n    sys.stdout.flush()\n\n\ndef format_time(seconds):\n    days = int(seconds / 3600 / 24)\n    seconds = seconds - days * 3600 * 24\n    hours = int(seconds / 3600)\n    seconds = seconds - hours * 3600\n    minutes = int(seconds / 60)\n    seconds = seconds - minutes * 60\n    secondsf = int(seconds)\n    seconds = seconds - secondsf\n    millis = int(seconds * 1000)\n\n    f = ''\n    i = 1\n    if days > 0:\n        f += str(days) + 'D'\n        i += 1\n    if hours > 0 and i <= 2:\n        f += str(hours) + 'h'\n        i += 1\n    if minutes > 0 and i <= 2:\n        f += str(minutes) + 'm'\n        i += 1\n    if secondsf > 0 and i <= 2:\n        f += str(secondsf) + 's'\n        i += 1\n    if millis > 0 and i <= 2:\n        f += str(millis) + 'ms'\n        i += 1\n    if f == '':\n        f = '0ms'\n    return f\n\n\ndef init_log(output_dir):\n    logging.basicConfig(level=logging.DEBUG,\n                        format='%(asctime)s %(message)s',\n                        datefmt='%Y%m%d-%H:%M:%S',\n                        filename=os.path.join(output_dir, 'log.log'),\n                        filemode='w')\n    console = logging.StreamHandler()\n    console.setLevel(logging.INFO)\n    logging.getLogger('').addHandler(console)\n    return logging","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:52:15.978760Z","iopub.execute_input":"2022-06-25T04:52:15.981265Z","iopub.status.idle":"2022-06-25T04:52:16.013629Z","shell.execute_reply.started":"2022-06-25T04:52:15.981226Z","shell.execute_reply":"2022-06-25T04:52:16.012831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 利用resnet网络得到特征图","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport math\nimport torch.utils.model_zoo as model_zoo\n\n__all__ = ['ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101',\n           'resnet152']\n\nmodel_urls = {\n    'resnet18': 'https://download.pytorch.org/models/resnet18-5c106cde.pth',\n    'resnet34': 'https://download.pytorch.org/models/resnet34-333f7ec4.pth',\n    'resnet50': 'https://download.pytorch.org/models/resnet50-19c8e357.pth',\n    'resnet101': 'https://download.pytorch.org/models/resnet101-5d3b4d8f.pth',\n    'resnet152': 'https://download.pytorch.org/models/resnet152-b121ed2d.pth',\n}\n\n\ndef conv3x3(in_planes, out_planes, stride=1):\n    \"3x3 convolution with padding\"\n    return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,\n                     padding=1, bias=False)\n\n\nclass BasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None):\n        super(BasicBlock, self).__init__()\n        self.conv1 = conv3x3(inplanes, planes, stride)\n        self.bn1 = nn.BatchNorm2d(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3(planes, planes)\n        self.bn2 = nn.BatchNorm2d(planes)\n        self.downsample = downsample\n        self.stride = stride\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        out = self.conv2(out)\n        out = self.bn2(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\n\n\nclass Bottleneck(nn.Module):\n    expansion = 4\n\n    def __init__(self, inplanes, planes, stride=1, downsample=None):\n        super(Bottleneck, self).__init__()\n        self.conv1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(planes)\n        self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, stride=stride,\n                               padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(planes)\n        self.conv3 = nn.Conv2d(planes, planes * 4, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm2d(planes * 4)\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n        self.stride = stride\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        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\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\n\n\nclass ResNet(nn.Module):\n    def __init__(self, block, layers, num_classes=1000):\n        self.inplanes = 64\n        super(ResNet, self).__init__()\n        self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3,\n                               bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU(inplace=True)\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.AvgPool2d(7)\n        self.fc = nn.Linear(512 * block.expansion, num_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels\n                m.weight.data.normal_(0, math.sqrt(2. / n))\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\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.Conv2d(self.inplanes, planes * block.expansion,\n                          kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(planes * block.expansion),\n            )\n\n        layers = []\n        layers.append(block(self.inplanes, planes, stride, downsample))\n        self.inplanes = planes * block.expansion\n        for i in range(1, blocks):\n            layers.append(block(self.inplanes, planes))\n\n        return nn.Sequential(*layers)\n\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        feature1 = x\n        x = self.avgpool(x)\n        x = x.view(x.size(0), -1)\n        x = nn.Dropout(p=0.5)(x)\n        feature2 = x\n        x = self.fc(x)\n\n        return x, feature1, feature2\n\n\ndef resnet18(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-18 model.\n\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(BasicBlock, [2, 2, 2, 2], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet18']))\n    return model\n\n\ndef resnet34(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-34 model.\n\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(BasicBlock, [3, 4, 6, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet34']))\n    return model\n\n\ndef resnet50(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-50 model.\n\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 4, 6, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet50']))\n    return model\n\n\ndef resnet101(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-101 model.\n\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 4, 23, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet101']))\n    return model\n\n\ndef resnet152(pretrained=False, **kwargs):\n    \"\"\"Constructs a ResNet-152 model.\n\n    Args:\n        pretrained (bool): If True, returns a model pre-trained on ImageNet\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 8, 36, 3], **kwargs)\n    if pretrained:\n        model.load_state_dict(model_zoo.load_url(model_urls['resnet152']))\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:52:16.018704Z","iopub.execute_input":"2022-06-25T04:52:16.021846Z","iopub.status.idle":"2022-06-25T04:52:16.089804Z","shell.execute_reply.started":"2022-06-25T04:52:16.021740Z","shell.execute_reply":"2022-06-25T04:52:16.089012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# NTS_Net网络","metadata":{}},{"cell_type":"code","source":"from torch import nn\nimport torch\nimport torch.nn.functional as F\nfrom torch.autograd import Variable\nimport numpy as np\n# from core.anchors import generate_default_anchor_maps, hard_nms\n# from config import CAT_NUM, PROPOSAL_NUM\n\n\nclass ProposalNet(nn.Module):\n    def __init__(self):\n        super(ProposalNet, self).__init__()\n        self.down1 = nn.Conv2d(2048, 128, 3, 1, 1)\n        self.down2 = nn.Conv2d(128, 128, 3, 2, 1)\n        self.down3 = nn.Conv2d(128, 128, 3, 2, 1)\n        self.ReLU = nn.ReLU()\n        self.tidy1 = nn.Conv2d(128, 6, 1, 1, 0)\n        self.tidy2 = nn.Conv2d(128, 6, 1, 1, 0)\n        self.tidy3 = nn.Conv2d(128, 9, 1, 1, 0)\n\n    def forward(self, x):\n        batch_size = x.size(0)\n        d1 = self.ReLU(self.down1(x))\n        d2 = self.ReLU(self.down2(d1))\n        d3 = self.ReLU(self.down3(d2))\n        t1 = self.tidy1(d1).view(batch_size, -1)\n        t2 = self.tidy2(d2).view(batch_size, -1)\n        t3 = self.tidy3(d3).view(batch_size, -1)\n        return torch.cat((t1, t2, t3), dim=1)\n\n\nclass attention_net(nn.Module):\n    def __init__(self, topN=4):\n        super(attention_net, self).__init__()\n        self.pretrained_model = resnet50(pretrained=True)\n        self.pretrained_model.avgpool = nn.AdaptiveAvgPool2d(1)\n#         self.pretrained_model.fc = nn.Linear(512 * 4, 200)\n        self.pretrained_model.fc = nn.Linear(512 * 4, 100)\n        self.proposal_net = ProposalNet()\n        self.topN = topN\n#         self.concat_net = nn.Linear(2048 * (CAT_NUM + 1), 200)\n#         self.partcls_net = nn.Linear(512 * 4, 200)\n        self.concat_net = nn.Linear(2048 * (CAT_NUM + 1), 100)\n        self.partcls_net = nn.Linear(512 * 4, 100)\n        _, edge_anchors, _ = generate_default_anchor_maps()\n        self.pad_side = 224\n        self.edge_anchors = (edge_anchors + 224).astype(np.int)\n\n    def forward(self, x):\n        resnet_out, rpn_feature, feature = self.pretrained_model(x)\n        x_pad = F.pad(x, (self.pad_side, self.pad_side, self.pad_side, self.pad_side), mode='constant', value=0)\n        batch = x.size(0)\n        # we will reshape rpn to shape: batch * nb_anchor\n        rpn_score = self.proposal_net(rpn_feature.detach())\n        all_cdds = [\n            np.concatenate((x.reshape(-1, 1), self.edge_anchors.copy(), np.arange(0, len(x)).reshape(-1, 1)), axis=1)\n            for x in rpn_score.data.cpu().numpy()]\n        top_n_cdds = [hard_nms(x, topn=self.topN, iou_thresh=0.25) for x in all_cdds]\n        top_n_cdds = np.array(top_n_cdds)\n        top_n_index = top_n_cdds[:, :, -1].astype(np.int)\n        top_n_index = torch.from_numpy(top_n_index).cuda()\n        top_n_prob = torch.gather(rpn_score, dim=1, index=top_n_index)\n        part_imgs = torch.zeros([batch, self.topN, 3, 224, 224]).cuda()\n        for i in range(batch):\n            for j in range(self.topN):\n                [y0, x0, y1, x1] = top_n_cdds[i][j, 1:5].astype(np.int)\n                part_imgs[i:i + 1, j] = F.interpolate(x_pad[i:i + 1, :, y0:y1, x0:x1], size=(224, 224), mode='bilinear',\n                                                      align_corners=True)\n        part_imgs = part_imgs.view(batch * self.topN, 3, 224, 224)\n        _, _, part_features = self.pretrained_model(part_imgs.detach())\n        part_feature = part_features.view(batch, self.topN, -1)\n        part_feature = part_feature[:, :CAT_NUM, ...].contiguous()\n        part_feature = part_feature.view(batch, -1)\n        # concat_logits have the shape: B*200\n        concat_out = torch.cat([part_feature, feature], dim=1)\n        concat_logits = self.concat_net(concat_out)\n        raw_logits = resnet_out\n        # part_logits have the shape: B*N*200\n        part_logits = self.partcls_net(part_features).view(batch, self.topN, -1)\n        return [raw_logits, concat_logits, part_logits, top_n_index, top_n_prob]\n\n\ndef list_loss(logits, targets):\n    temp = F.log_softmax(logits, -1)\n    loss = [-temp[i][targets[i].item()] for i in range(logits.size(0))]\n    return torch.stack(loss)\n\n\ndef ranking_loss(score, targets, proposal_num=PROPOSAL_NUM):\n    loss = Variable(torch.zeros(1).cuda())\n    batch_size = score.size(0)\n    for i in range(proposal_num):\n        targets_p = (targets > targets[:, i].unsqueeze(1)).type(torch.cuda.FloatTensor)\n        pivot = score[:, i].unsqueeze(1)\n        loss_p = (1 - pivot + score) * targets_p\n        loss_p = torch.sum(F.relu(loss_p))\n        loss += loss_p\n    return loss / batch_size","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:52:16.094232Z","iopub.execute_input":"2022-06-25T04:52:16.097116Z","iopub.status.idle":"2022-06-25T04:52:16.146455Z","shell.execute_reply.started":"2022-06-25T04:52:16.097074Z","shell.execute_reply":"2022-06-25T04:52:16.145868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"import os\nimport torch.utils.data\nfrom torch.nn import DataParallel\nfrom datetime import datetime\nfrom torch.optim.lr_scheduler import MultiStepLR\n# from config import BATCH_SIZE, PROPOSAL_NUM, SAVE_FREQ, LR, WD, resume, save_dir\n# from core import model, dataset\n# from core.utils import init_log, progress_bar\n\ntest_acc_set=[]#测试集的准确率\ntrain_loss_set=[]#训练集的损失\nos.environ['CUDA_VISIBLE_DEVICES'] = '0,1,2,3'\nstart_epoch = 1\n# save_dir = os.path.join(save_dir, datetime.now().strftime('%Y%m%d_%H%M%S'))\n# if os.path.exists(save_dir):\n#     raise NameError('model dir exists!')\n# os.makedirs(save_dir)\nlogging = init_log(save_dir)\n_print = logging.info\n\n# read dataset\n# trainset = dataset.CUB(root='./CUB_200_2011', is_train=True, data_len=None)\n# trainloader = torch.utils.data.DataLoader(trainset, batch_size=BATCH_SIZE,\n#                                           shuffle=True, num_workers=8, drop_last=False)\n# testset = dataset.CUB(root='./CUB_200_2011', is_train=False, data_len=None)\n# testloader = torch.utils.data.DataLoader(testset, batch_size=BATCH_SIZE,\n#                                          shuffle=False, num_workers=8, drop_last=False)\n\ntrainset=MyDataset(train_msg, mode='train')\ntrainloader = torch.utils.data.DataLoader(trainset, batch_size=BATCH_SIZE,\n                                          shuffle=True, num_workers=8, drop_last=False)\nvalset=MyDataset(train_msg, mode='val')\ntestloader = torch.utils.data.DataLoader(valset, batch_size=BATCH_SIZE,\n                                         shuffle=False, num_workers=8, drop_last=False)\n\n# define model\n# net = model.attention_net(topN=PROPOSAL_NUM)\nnet = attention_net(topN=PROPOSAL_NUM)\nif resume:\n    ckpt = torch.load(resume)\n    net.load_state_dict(ckpt['net_state_dict'])\n    start_epoch = ckpt['epoch'] + 1\ncreterion = torch.nn.CrossEntropyLoss()\n\n# define optimizers\nraw_parameters = list(net.pretrained_model.parameters())\npart_parameters = list(net.proposal_net.parameters())\nconcat_parameters = list(net.concat_net.parameters())\npartcls_parameters = list(net.partcls_net.parameters())\n\nraw_optimizer = torch.optim.SGD(raw_parameters, lr=LR, momentum=0.9, weight_decay=WD)\nconcat_optimizer = torch.optim.SGD(concat_parameters, lr=LR, momentum=0.9, weight_decay=WD)\npart_optimizer = torch.optim.SGD(part_parameters, lr=LR, momentum=0.9, weight_decay=WD)\npartcls_optimizer = torch.optim.SGD(partcls_parameters, lr=LR, momentum=0.9, weight_decay=WD)\nschedulers = [MultiStepLR(raw_optimizer, milestones=[60, 100], gamma=0.1),\n              MultiStepLR(concat_optimizer, milestones=[60, 100], gamma=0.1),\n              MultiStepLR(part_optimizer, milestones=[60, 100], gamma=0.1),\n              MultiStepLR(partcls_optimizer, milestones=[60, 100], gamma=0.1)]\nnet = net.cuda()\nnet = DataParallel(net)\n\n# for epoch in range(start_epoch, 500):\nfor epoch in range(start_epoch, 70):\n    for scheduler in schedulers:\n        scheduler.step()\n\n    # begin training\n    _print('--' * 50)\n    net.train()\n    train_loss = 0\n    train_correct = 0\n    total = 0\n    for i, data in enumerate(trainloader):\n        img, label = data[0].cuda(), data[1].cuda()\n        batch_size = img.size(0)\n        raw_optimizer.zero_grad()\n        part_optimizer.zero_grad()\n        concat_optimizer.zero_grad()\n        partcls_optimizer.zero_grad()\n\n        raw_logits, concat_logits, part_logits, _, top_n_prob = net(img)\n        part_loss = list_loss(part_logits.view(batch_size * PROPOSAL_NUM, -1),\n                                    label.unsqueeze(1).repeat(1, PROPOSAL_NUM).view(-1)).view(batch_size, PROPOSAL_NUM)\n        raw_loss = creterion(raw_logits, label)\n        concat_loss = creterion(concat_logits, label)\n        _, concat_predict = torch.max(concat_logits, 1)\n        train_correct += torch.sum(concat_predict.data == label.data)\n        rank_loss = ranking_loss(top_n_prob, part_loss)\n        partcls_loss = creterion(part_logits.view(batch_size * PROPOSAL_NUM, -1),\n                                 label.unsqueeze(1).repeat(1, PROPOSAL_NUM).view(-1))\n        total+=batch_size\n        train_loss+=concat_loss.item() * batch_size\n\n        total_loss = raw_loss + rank_loss + concat_loss + partcls_loss\n        total_loss.backward()\n        raw_optimizer.step()\n        part_optimizer.step()\n        concat_optimizer.step()\n        partcls_optimizer.step()\n        progress_bar(i, len(trainloader), 'train')\n        \n      \n    train_loss=train_loss/total\n    train_acc = float(train_correct) / total \n    train_loss_set.append(train_loss)\n    print(\n            'epoch:{} - train loss: {:.3f} and train_acc: {:.3f} total sample: {}'.format(\n                epoch,\n                train_loss,\n                train_acc,\n                total))\n        \n\n    if epoch % SAVE_FREQ == 0:\n#         train_loss = 0\n#         train_correct = 0\n#         total = 0\n#         net.eval()\n#         for i, data in enumerate(trainloader):\n#             with torch.no_grad():\n#                 img, label = data[0].cuda(), data[1].cuda()\n#                 batch_size = img.size(0)\n#                 _, concat_logits, _, _, _ = net(img)\n#                 # calculate loss\n#                 concat_loss = creterion(concat_logits, label)\n#                 # calculate accuracy\n#                 _, concat_predict = torch.max(concat_logits, 1)\n#                 total += batch_size\n#                 train_correct += torch.sum(concat_predict.data == label.data)\n#                 train_loss += concat_loss.item() * batch_size\n#                 progress_bar(i, len(trainloader), 'eval train set')\n\n#         train_acc = float(train_correct) / total\n#         train_loss = train_loss / total\n\n#         print(\n#             'epoch:{} - train loss: {:.3f} and train acc: {:.3f} total sample: {}'.format(\n#                 epoch,\n#                 train_loss,\n#                 train_acc,\n#                 total))\n\n\t# evaluate on val set\n        test_loss = 0\n        test_correct = 0\n        total = 0\n        for i, data in enumerate(testloader):\n            with torch.no_grad():\n                img, label = data[0].cuda(), data[1].cuda()\n                batch_size = img.size(0)\n                _, concat_logits, _, _, _ = net(img)\n                # calculate loss\n                concat_loss = creterion(concat_logits, label)\n                # calculate accuracy\n                _, concat_predict = torch.max(concat_logits, 1)\n                total += batch_size\n                test_correct += torch.sum(concat_predict.data == label.data)\n                test_loss += concat_loss.item() * batch_size\n                progress_bar(i, len(testloader), 'eval test set')\n\n        test_acc = float(test_correct) / total\n        test_loss = test_loss / total\n        test_acc_set.append(test_acc)\n        print(\n            'epoch:{} - test loss: {:.3f} and test acc: {:.3f} total sample: {}'.format(\n                epoch,\n                test_loss,\n                test_acc,\n                total))\n\n\t# save model\n        net_state_dict = net.module.state_dict()\n        if not os.path.exists(save_dir):\n            os.makedirs(save_dir)\n        torch.save({\n            'epoch': epoch,\n            'train_loss': train_loss,\n            'train_acc': train_acc,\n            'test_loss': test_loss,\n            'test_acc': test_acc,\n            'net_state_dict': net_state_dict},\n            os.path.join(save_dir, '%03d.ckpt' % epoch))\n        \n        print(train_loss_set)\n        print(test_acc_set)\n\nprint('finishing training')","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:58:40.684869Z","iopub.execute_input":"2022-06-25T04:58:40.685147Z","iopub.status.idle":"2022-06-25T13:26:15.217189Z","shell.execute_reply.started":"2022-06-25T04:58:40.685117Z","shell.execute_reply":"2022-06-25T13:26:15.213941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# test","metadata":{}},{"cell_type":"code","source":"#加载测试数据集\ntest_msg=pd.read_csv(os.path.join(root,'sample_submission.csv'))\ntest_msg.dropna(inplace=True) \nprint (f\"number of test data = {len(test_msg)}\")\ntest_msg.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:53:26.931338Z","iopub.status.idle":"2022-06-25T04:53:26.931900Z","shell.execute_reply.started":"2022-06-25T04:53:26.931662Z","shell.execute_reply":"2022-06-25T04:53:26.931688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_msg[\"dir\"] = test_msg[\"filename\"].apply( lambda image:IMG_TEST+image )\ntest_msg[\"dir_is_true\"] = test_msg[\"dir\"].apply(lambda dir: os.path.exists(dir))\ntest_msg[\"sorghum_type\"] = test_msg[\"cultivar\"].map(lambda type:sorghum_type.index(type))\ntest_msg.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:53:26.933045Z","iopub.status.idle":"2022-06-25T04:53:26.933610Z","shell.execute_reply.started":"2022-06-25T04:53:26.933377Z","shell.execute_reply":"2022-06-25T04:53:26.933403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testset=TestDataset(test_msg)\ntloader = torch.utils.data.DataLoader(testset, batch_size=BATCH_SIZE,\n                                         shuffle=False, num_workers=8, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:53:26.934732Z","iopub.status.idle":"2022-06-25T04:53:26.935345Z","shell.execute_reply.started":"2022-06-25T04:53:26.935098Z","shell.execute_reply":"2022-06-25T04:53:26.935122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TTA","metadata":{}},{"cell_type":"code","source":"# base\n\nimport itertools\nfrom functools import partial\nfrom typing import List, Optional, Union\n\n# from . import functional as F\n\n\nclass BaseTransform:\n    identity_param = None\n\n    def __init__(\n            self,\n            name: str,\n            params: Union[list, tuple],\n    ):\n        self.params = params\n        self.pname = name\n\n    def apply_aug_image(self, image, *args, **params):\n        raise NotImplementedError\n\n    def apply_deaug_mask(self, mask, *args, **params):\n        raise NotImplementedError\n\n    def apply_deaug_label(self, label, *args, **params):\n        raise NotImplementedError\n\n    def apply_deaug_keypoints(self, keypoints, *args, **params):\n        raise NotImplementedError\n\n\nclass ImageOnlyTransform(BaseTransform):\n\n    def apply_deaug_mask(self, mask, *args, **params):\n        return mask\n\n    def apply_deaug_label(self, label, *args, **params):\n        return label\n\n    def apply_deaug_keypoints(self, keypoints, *args, **params):\n        return keypoints\n\n\nclass DualTransform(BaseTransform):\n    pass\n\n\nclass Chain:\n\n    def __init__(\n            self,\n            functions: List[callable]\n    ):\n        self.functions = functions or []\n\n    def __call__(self, x):\n        for f in self.functions:\n            x = f(x)\n        return x\n\n\nclass Transformer:\n    def __init__(\n            self,\n            image_pipeline: Chain,\n            mask_pipeline: Chain,\n            label_pipeline: Chain,\n            keypoints_pipeline: Chain\n    ):\n        self.image_pipeline = image_pipeline\n        self.mask_pipeline = mask_pipeline\n        self.label_pipeline = label_pipeline\n        self.keypoints_pipeline = keypoints_pipeline\n\n    def augment_image(self, image):\n        return self.image_pipeline(image)\n\n    def deaugment_mask(self, mask):\n        return self.mask_pipeline(mask)\n\n    def deaugment_label(self, label):\n        return self.label_pipeline(label)\n\n    def deaugment_keypoints(self, keypoints):\n        return self.keypoints_pipeline(keypoints)\n\n\nclass Compose:\n\n    def __init__(\n            self,\n            transforms: List[BaseTransform],\n    ):\n        self.aug_transforms = transforms\n        self.aug_transform_parameters = list(itertools.product(*[t.params for t in self.aug_transforms]))\n        self.deaug_transforms = transforms[::-1]\n        self.deaug_transform_parameters = [p[::-1] for p in self.aug_transform_parameters]\n\n    def __iter__(self) -> Transformer:\n        for aug_params, deaug_params in zip(self.aug_transform_parameters, self.deaug_transform_parameters):\n            image_aug_chain = Chain([partial(t.apply_aug_image, **{t.pname: p})\n                                     for t, p in zip(self.aug_transforms, aug_params)])\n            mask_deaug_chain = Chain([partial(t.apply_deaug_mask, **{t.pname: p})\n                                      for t, p in zip(self.deaug_transforms, deaug_params)])\n            label_deaug_chain = Chain([partial(t.apply_deaug_label, **{t.pname: p})\n                                       for t, p in zip(self.deaug_transforms, deaug_params)])\n            keypoints_deaug_chain = Chain([partial(t.apply_deaug_keypoints, **{t.pname: p})\n                                           for t, p in zip(self.deaug_transforms, deaug_params)])\n            yield Transformer(\n                image_pipeline=image_aug_chain,\n                mask_pipeline=mask_deaug_chain,\n                label_pipeline=label_deaug_chain,\n                keypoints_pipeline=keypoints_deaug_chain\n            )\n\n    def __len__(self) -> int:\n        return len(self.aug_transform_parameters)\n\n\nclass Merger:\n\n    def __init__(\n            self,\n            type: str = 'mean',\n            n: int = 1,\n    ):\n\n        if type not in ['mean', 'gmean', 'sum', 'max', 'min', 'tsharpen']:\n            raise ValueError('Not correct merge type `{}`.'.format(type))\n\n        self.output = None\n        self.type = type\n        self.n = n\n\n    def append(self, x):\n\n        if self.type == 'tsharpen':\n            x = x ** 0.5\n\n        if self.output is None:\n            self.output = x\n        elif self.type in ['mean', 'sum', 'tsharpen']:\n            self.output = self.output + x\n        elif self.type == 'gmean':\n            self.output = self.output * x\n        elif self.type == 'max':\n            self.output = F.max(self.output, x)\n        elif self.type == 'min':\n            self.output = F.min(self.output, x)\n\n    @property\n    def result(self):\n        if self.type in ['sum', 'max', 'min']:\n            result = self.output\n        elif self.type in ['mean', 'tsharpen']:\n            result = self.output / self.n\n        elif self.type in ['gmean']:\n            result = self.output ** (1 / self.n)\n        else:\n            raise ValueError('Not correct merge type `{}`.'.format(self.type))\n        return result\n","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:53:26.936516Z","iopub.status.idle":"2022-06-25T04:53:26.937077Z","shell.execute_reply.started":"2022-06-25T04:53:26.936826Z","shell.execute_reply":"2022-06-25T04:53:26.936852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TTA\n\nfrom typing import Optional, Mapping, Union, Tuple\n\nclass ClassificationTTAWrapper(nn.Module):\n    \"\"\"Wrap PyTorch nn.Module (classification model) with test time augmentation transforms\n\n    Args:\n        model (torch.nn.Module): classification model with single input and single output\n            (.forward(x) should return either torch.Tensor or Mapping[str, torch.Tensor])\n        transforms (ttach.Compose): composition of test time transforms\n        merge_mode (str): method to merge augmented predictions mean/gmean/max/min/sum/tsharpen\n        output_label_key (str): if model output is `dict`, specify which key belong to `label`\n    \"\"\"\n\n    def __init__(\n        self,\n        model: nn.Module,\n        transforms: Compose,\n        merge_mode: str = \"mean\",\n        # output_label_key: Optional[str] = None,\n        output_label_key: Optional[int] = None,\n    ):\n        super().__init__()\n        self.model = model\n        self.transforms = transforms\n        self.merge_mode = merge_mode\n        self.output_key = output_label_key\n\n    def forward(\n        self, image: torch.Tensor, *args\n    ) -> Union[torch.Tensor, Mapping[str, torch.Tensor]]:\n        merger = Merger(type=self.merge_mode, n=len(self.transforms))\n\n        for transformer in self.transforms:\n            augmented_image = transformer.augment_image(image)\n            augmented_output = self.model(augmented_image, *args)\n            if self.output_key is not None:\n                augmented_output = augmented_output[self.output_key]\n            deaugmented_output = transformer.deaugment_label(augmented_output)\n            merger.append(deaugmented_output)\n\n        result = merger.result\n#         if self.output_key is not None:\n# #             result = {self.output_key: result}\n#             result = {self.output_key: result}\n\n        return result\n","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:53:26.938191Z","iopub.status.idle":"2022-06-25T04:53:26.938733Z","shell.execute_reply.started":"2022-06-25T04:53:26.938498Z","shell.execute_reply":"2022-06-25T04:53:26.938524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 加载模型并测试","metadata":{}},{"cell_type":"code","source":"def rot90(x, k=1):\n    \"\"\"rotate batch of images by 90 degrees k times\"\"\"\n    return torch.rot90(x, k, (2, 3))\n\n\ndef hflip(x):\n    \"\"\"flip batch of images horizontally\"\"\"\n    return x.flip(3)\n\n\ndef vflip(x):\n    \"\"\"flip batch of images vertically\"\"\"\n    return x.flip(2)\n\n\ndef sum(x1, x2):\n    \"\"\"sum of two tensors\"\"\"\n    return x1 + x2\n\n\ndef add(x, value):\n    \"\"\"add value to tensor\"\"\"\n    return x + value\n\n\ndef max(x1, x2):\n    \"\"\"compare 2 tensors and take max values\"\"\"\n    return torch.max(x1, x2)\n\n\ndef min(x1, x2):\n    \"\"\"compare 2 tensors and take min values\"\"\"\n    return torch.min(x1, x2)\n\n\ndef multiply(x, factor):\n    \"\"\"multiply tensor by factor\"\"\"\n    return x * factor\n\n\ndef scale(x, scale_factor, interpolation=\"nearest\", align_corners=None):\n    \"\"\"scale batch of images by `scale_factor` with given interpolation mode\"\"\"\n    h, w = x.shape[2:]\n    new_h = int(h * scale_factor)\n    new_w = int(w * scale_factor)\n    return F.interpolate(\n        x, size=(new_h, new_w), mode=interpolation, align_corners=align_corners\n    )\n\n\ndef resize(x, size, interpolation=\"nearest\", align_corners=None):\n    \"\"\"resize batch of images to given spatial size with given interpolation mode\"\"\"\n    return F.interpolate(x, size=size, mode=interpolation, align_corners=align_corners)\n\ndef crop(x, x_min=None, x_max=None, y_min=None, y_max=None):\n    \"\"\"perform crop on batch of images\"\"\"\n    return x[:, :, y_min:y_max, x_min:x_max]\n\ndef crop_lt(x, crop_h, crop_w):\n    \"\"\"crop left top corner\"\"\"\n    return x[:, :, 0:crop_h, 0:crop_w]\n\n\ndef crop_lb(x, crop_h, crop_w):\n    \"\"\"crop left bottom corner\"\"\"\n    return x[:, :, -crop_h:, 0:crop_w]\n\n\ndef crop_rt(x, crop_h, crop_w):\n    \"\"\"crop right top corner\"\"\"\n    return x[:, :, 0:crop_h, -crop_w:]\n\n\ndef crop_rb(x, crop_h, crop_w):\n    \"\"\"crop right bottom corner\"\"\"\n    return x[:, :, -crop_h:, -crop_w:]\n\n\ndef center_crop(x, crop_h, crop_w):\n    \"\"\"make center crop\"\"\"\n\n    center_h = x.shape[2] // 2\n    center_w = x.shape[3] // 2\n    half_crop_h = crop_h // 2\n    half_crop_w = crop_w // 2\n\n    y_min = center_h - half_crop_h\n    y_max = center_h + half_crop_h + crop_h % 2\n    x_min = center_w - half_crop_w\n    x_max = center_w + half_crop_w + crop_w % 2\n\n    return x[:, :, y_min:y_max, x_min:x_max]\n\ndef _disassemble_keypoints(keypoints):\n    x = keypoints[..., 0]\n    y = keypoints[..., 1]\n    return x, y\n\n\ndef _assemble_keypoints(x, y):\n    return torch.stack([x, y], dim=-1)\n\n\ndef keypoints_hflip(keypoints):\n    x, y = _disassemble_keypoints(keypoints)\n    return _assemble_keypoints(1. - x, y)\n\n\ndef keypoints_vflip(keypoints):\n    x, y = _disassemble_keypoints(keypoints)\n    return _assemble_keypoints(x, 1. - y)\n\n\ndef keypoints_rot90(keypoints, k=1):\n\n    if k not in {0, 1, 2, 3}:\n        raise ValueError(\"Parameter k must be in [0:3]\")\n    if k == 0:\n        return keypoints\n    x, y = _disassemble_keypoints(keypoints)\n\n    if k == 1:\n        xy = [y, 1. - x]\n    elif k == 2:\n        xy = [1. - x, 1. - y]\n    elif k == 3:\n        xy = [1. - y, x]\n\n    return _assemble_keypoints(*xy)\n\n\n# 各种图像变换\nclass HorizontalFlip(DualTransform):\n    \"\"\"Flip images horizontally (left->right)\"\"\"\n\n    identity_param = False\n\n    def __init__(self):\n        super().__init__(\"apply\", [False, True])\n\n    def apply_aug_image(self, image, apply=False, **kwargs):\n        if apply:\n            image = hflip(image)\n        return image\n\n    def apply_deaug_mask(self, mask, apply=False, **kwargs):\n        if apply:\n            mask = hflip(mask)\n        return mask\n\n    def apply_deaug_label(self, label, apply=False, **kwargs):\n        return label\n\n    def apply_deaug_keypoints(self, keypoints, apply=False, **kwargs):\n        if apply:\n            keypoints = keypoints_hflip(keypoints)\n        return keypoints\n\n\nclass VerticalFlip(DualTransform):\n    \"\"\"Flip images vertically (up->down)\"\"\"\n\n    identity_param = False\n\n    def __init__(self):\n        super().__init__(\"apply\", [False, True])\n\n    def apply_aug_image(self, image, apply=False, **kwargs):\n        if apply:\n            image = vflip(image)\n        return image\n\n    def apply_deaug_mask(self, mask, apply=False, **kwargs):\n        if apply:\n            mask = vflip(mask)\n        return mask\n\n    def apply_deaug_label(self, label, apply=False, **kwargs):\n        return label\n\n    def apply_deaug_keypoints(self, keypoints, apply=False, **kwargs):\n        if apply:\n            keypoints = keypoints_vflip(keypoints)\n        return keypoints\n\n\nclass Rotate90(DualTransform):\n    \"\"\"Rotate images 0/90/180/270 degrees\n\n    Args:\n        angles (list): angles to rotate images\n    \"\"\"\n\n    identity_param = 0\n\n    def __init__(self, angles: List[int]):\n        if self.identity_param not in angles:\n            angles = [self.identity_param] + list(angles)\n\n        super().__init__(\"angle\", angles)\n\n    def apply_aug_image(self, image, angle=0, **kwargs):\n        k = angle // 90 if angle >= 0 else (angle + 360) // 90\n        return rot90(image, k)\n\n    def apply_deaug_mask(self, mask, angle=0, **kwargs):\n        return self.apply_aug_image(mask, -angle)\n\n    def apply_deaug_label(self, label, angle=0, **kwargs):\n        return label\n\n    def apply_deaug_keypoints(self, keypoints, angle=0, **kwargs):\n        angle *= -1\n        k = angle // 90 if angle >= 0 else (angle + 360) // 90\n        return keypoints_rot90(keypoints, k=k)\n\n\nclass Scale(DualTransform):\n    \"\"\"Scale images\n\n    Args:\n        scales (List[Union[int, float]]): scale factors for spatial image dimensions\n        interpolation (str): one of \"nearest\"/\"lenear\" (see more in torch.nn.interpolate)\n        align_corners (bool): see more in torch.nn.interpolate\n    \"\"\"\n\n    identity_param = 1\n\n    def __init__(\n        self,\n        scales: List[Union[int, float]],\n        interpolation: str = \"nearest\",\n        align_corners: Optional[bool] = None,\n    ):\n        if self.identity_param not in scales:\n            scales = [self.identity_param] + list(scales)\n        self.interpolation = interpolation\n        self.align_corners = align_corners\n\n        super().__init__(\"scale\", scales)\n\n    def apply_aug_image(self, image, scale=1, **kwargs):\n        if scale != self.identity_param:\n            image = scale(\n                image,\n                scale,\n                interpolation=self.interpolation,\n                align_corners=self.align_corners,\n            )\n        return image\n\n    def apply_deaug_mask(self, mask, scale=1, **kwargs):\n        if scale != self.identity_param:\n            mask = scale(\n                mask,\n                1 / scale,\n                interpolation=self.interpolation,\n                align_corners=self.align_corners,\n            )\n        return mask\n\n    def apply_deaug_label(self, label, scale=1, **kwargs):\n        return label\n\n    def apply_deaug_keypoints(self, keypoints, scale=1, **kwargs):\n        return keypoints\n\n\nclass Resize(DualTransform):\n    \"\"\"Resize images\n\n    Args:\n        sizes (List[Tuple[int, int]): scale factors for spatial image dimensions\n        original_size Tuple(int, int): optional, image original size for deaugmenting mask\n        interpolation (str): one of \"nearest\"/\"lenear\" (see more in torch.nn.interpolate)\n        align_corners (bool): see more in torch.nn.interpolate\n    \"\"\"\n\n    def __init__(\n        self,\n        sizes: List[Tuple[int, int]],\n        original_size: Tuple[int, int] = None,\n        interpolation: str = \"nearest\",\n        align_corners: Optional[bool] = None,\n    ):\n        if original_size is not None and original_size not in sizes:\n            sizes = [original_size] + list(sizes)\n        self.interpolation = interpolation\n        self.align_corners = align_corners\n        self.original_size = original_size\n\n        super().__init__(\"size\", sizes)\n\n    def apply_aug_image(self, image, size, **kwargs):\n        if size != self.original_size:\n            image = resize(\n                image,\n                size,\n                interpolation=self.interpolation,\n                align_corners=self.align_corners,\n            )\n        return image\n\n    def apply_deaug_mask(self, mask, size, **kwargs):\n        if self.original_size is None:\n            raise ValueError(\n                \"Provide original image size to make mask backward transformation\"\n            )\n        if size != self.original_size:\n            mask = resize(\n                mask,\n                self.original_size,\n                interpolation=self.interpolation,\n                align_corners=self.align_corners,\n            )\n        return mask\n\n    def apply_deaug_label(self, label, size=1, **kwargs):\n        return label\n\n    def apply_deaug_keypoints(self, keypoints, size=1, **kwargs):\n        return keypoints\n\n\nclass Add(ImageOnlyTransform):\n    \"\"\"Add value to images\n\n    Args:\n        values (List[float]): values to add to each pixel\n    \"\"\"\n\n    identity_param = 0\n\n    def __init__(self, values: List[float]):\n\n        if self.identity_param not in values:\n            values = [self.identity_param] + list(values)\n        super().__init__(\"value\", values)\n\n    def apply_aug_image(self, image, value=0, **kwargs):\n        if value != self.identity_param:\n            image = add(image, value)\n        return image\n\n\nclass Multiply(ImageOnlyTransform):\n    \"\"\"Multiply images by factor\n\n    Args:\n        factors (List[float]): factor to multiply each pixel by\n    \"\"\"\n\n    identity_param = 1\n\n    def __init__(self, factors: List[float]):\n        if self.identity_param not in factors:\n            factors = [self.identity_param] + list(factors)\n        super().__init__(\"factor\", factors)\n\n    def apply_aug_image(self, image, factor=1, **kwargs):\n        if factor != self.identity_param:\n            image = multiply(image, factor)\n        return image\n\n\n\nclass FiveCrops(ImageOnlyTransform):\n    \"\"\"Makes 4 crops for each corner + center crop\n\n    Args:\n        crop_height (int): crop height in pixels\n        crop_width (int): crop width in pixels \n    \"\"\"\n\n    def __init__(self, crop_height, crop_width):\n        crop_functions = (\n            partial(crop_lt, crop_h=crop_height, crop_w=crop_width),\n            partial(crop_lb, crop_h=crop_height, crop_w=crop_width),\n            partial(crop_rb, crop_h=crop_height, crop_w=crop_width),\n            partial(crop_rt, crop_h=crop_height, crop_w=crop_width),\n            partial(center_crop, crop_h=crop_height, crop_w=crop_width),\n        )\n        super().__init__(\"crop_fn\", crop_functions)\n\n    def apply_aug_image(self, image, crop_fn=None, **kwargs):\n        return crop_fn(image)\n\n    def apply_deaug_mask(self, mask, **kwargs):\n        raise ValueError(\"`FiveCrop` augmentation is not suitable for mask!\")\n\n    def apply_deaug_keypoints(self, keypoints, **kwargs):\n        raise ValueError(\"`FiveCrop` augmentation is not suitable for keypoints!\")\n\n\ndef five_crop_transform(crop_height, crop_width):\n    return Compose([FiveCrops(crop_height, crop_width)])\n\ndef ten_crop_transform(crop_height, crop_width):\n    return Compose([HorizontalFlip(), FiveCrops(crop_height, crop_width)])\n\ndef d4_transform():\n    return Compose(\n        [\n            HorizontalFlip(),\n            Rotate90(angles=[0, 90, 180, 270]),\n        ]\n    )\n\n","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:53:26.940082Z","iopub.status.idle":"2022-06-25T04:53:26.940629Z","shell.execute_reply.started":"2022-06-25T04:53:26.940399Z","shell.execute_reply":"2022-06-25T04:53:26.940424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.nn import DataParallel\n\ntest_model='../input/models4-nts/055.ckpt'\nnet = attention_net(topN=PROPOSAL_NUM)\nckpt = torch.load(test_model)\nnet.load_state_dict(ckpt['net_state_dict'])\nnet = net.cuda()\nnet = DataParallel(net)\ntta_model = ClassificationTTAWrapper(net, five_crop_transform(448,448), output_label_key = 1)\n# tta_model = ClassificationTTAWrapper(net, ten_crop_transform(448, 448), output_label_key = 1)\ncreterion = torch.nn.CrossEntropyLoss()\npredictions=[]\n\nfor i, data in enumerate(tloader):\n    with torch.no_grad():\n        img, label = data[0].cuda(), data[1].cuda()\n        batch_size = img.size(0)\n#         _, concat_logits, _, _, _ = net(img)\n#         concat_logits = tta_model.forward(img)\n        concat_logits = tta_model(img)\n        # calculate loss\n#         concat_loss = creterion(concat_logits, label)\n        # calculate accuracy\n        _, concat_predict = torch.max(concat_logits, 1)\n        predictions+=concat_predict.tolist()[:8]\n#         print(predictions)\n#        print(concat_predict.tolist()[:8])\n#         total += batch_size\n#         test_correct += torch.sum(concat_predict.data == label.data)\n#         test_loss += concat_loss.item() * batch_size\n        progress_bar(i, len(tloader), 'eval test set')\nsorghum_name = [sorghum_type[pred] for pred in predictions]\nsub = pd.read_csv(os.path.join(root , \"sample_submission.csv\"))\nsub[\"cultivar\"] = sorghum_name\nsub.to_csv('submission.csv', index=False)\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T04:53:26.941710Z","iopub.status.idle":"2022-06-25T04:53:26.942268Z","shell.execute_reply.started":"2022-06-25T04:53:26.942018Z","shell.execute_reply":"2022-06-25T04:53:26.942044Z"},"trusted":true},"execution_count":null,"outputs":[]}]}